""" Subtitle 字幕领域模型单元测试 """ import pytest from packages.domain.subtitle import ( SubtitleSegment, SubtitleTimeline, SubtitleWord, ) class TestSubtitleWord: """SubtitleWord 测试""" def test_duration_positive(self): word = SubtitleWord(text="你好", start=1.0, end=2.5) assert word.duration == pytest.approx(1.5) def test_duration_zero(self): word = SubtitleWord(text="a", start=5.0, end=5.0) assert word.duration == 0.0 def test_duration_negative_returns_zero(self): """测试结束时间小于开始时间时返回 0""" word = SubtitleWord(text="a", start=3.0, end=1.0) assert word.duration == 0.0 class TestSubtitleSegment: """SubtitleSegment 测试""" def test_duration(self): seg = SubtitleSegment(text="你好世界", start=0.0, end=3.0) assert seg.duration == pytest.approx(3.0) def test_duration_zero(self): seg = SubtitleSegment(text="test", start=5.0, end=5.0) assert seg.duration == 0.0 def test_duration_negative_returns_zero(self): seg = SubtitleSegment(text="test", start=5.0, end=2.0) assert seg.duration == 0.0 def test_char_count(self): seg = SubtitleSegment(text="你好世界", start=0, end=1) assert seg.char_count == 4 def test_char_count_empty(self): seg = SubtitleSegment(text="", start=0, end=1) assert seg.char_count == 0 def test_default_words_empty(self): seg = SubtitleSegment(text="test", start=0, end=1) assert seg.words == [] def test_with_words(self): words = [ SubtitleWord(text="你好", start=0.0, end=1.0), SubtitleWord(text="世界", start=1.0, end=2.0), ] seg = SubtitleSegment(text="你好世界", start=0.0, end=2.0, words=words) assert len(seg.words) == 2 assert seg.words[0].text == "你好" assert seg.words[1].text == "世界" class TestSubtitleTimelineBasics: """SubtitleTimeline 基础属性测试""" def test_empty_timeline(self): tl = SubtitleTimeline() assert tl.segment_count == 0 assert tl.total_chars == 0 assert tl.language == "zh" assert tl.total_duration == 0.0 def test_segment_count(self): tl = SubtitleTimeline( segments=[ SubtitleSegment(text="a", start=0, end=1), SubtitleSegment(text="b", start=1, end=2), SubtitleSegment(text="c", start=2, end=3), ] ) assert tl.segment_count == 3 def test_total_chars(self): tl = SubtitleTimeline( segments=[ SubtitleSegment(text="你好", start=0, end=1), SubtitleSegment(text="世界", start=1, end=2), SubtitleSegment(text="abcde", start=2, end=3), ] ) assert tl.total_chars == 9 def test_custom_language(self): tl = SubtitleTimeline(language="en") assert tl.language == "en" def test_custom_total_duration(self): tl = SubtitleTimeline(total_duration=60.0) assert tl.total_duration == 60.0 class TestMergeShortSegments: """merge_short_segments 测试""" def test_single_segment_no_merge(self): """单个片段不需要合并""" tl = SubtitleTimeline( segments=[ SubtitleSegment(text="a", start=0, end=1), ] ) result = tl.merge_short_segments(min_chars=8) assert result.segment_count == 1 assert result.segments[0].text == "a" def test_empty_timeline(self): """空时间轴""" tl = SubtitleTimeline() result = tl.merge_short_segments(min_chars=8) assert result.segment_count == 0 def test_all_short_segments_merge_into_one(self): """所有短片段合并成一个""" tl = SubtitleTimeline( segments=[ SubtitleSegment(text="你", start=0, end=0.5), SubtitleSegment(text="好", start=0.5, end=1.0), SubtitleSegment(text="世", start=1.0, end=1.5), SubtitleSegment(text="界", start=1.5, end=2.0), ] ) result = tl.merge_short_segments(min_chars=8) assert result.segment_count == 1 assert result.segments[0].text == "你好世界" assert result.segments[0].start == 0 assert result.segments[0].end == 2.0 def test_merge_short_segments_preserves_timing(self): """合并后时间轴正确""" tl = SubtitleTimeline( segments=[ SubtitleSegment(text="你好", start=1.0, end=2.0), SubtitleSegment(text="世界", start=2.0, end=3.5), ] ) result = tl.merge_short_segments(min_chars=10) assert result.segment_count == 1 assert result.segments[0].start == 1.0 assert result.segments[0].end == 3.5 def test_merge_short_segments_with_words(self): """合并后词级信息保留""" w1 = SubtitleWord(text="你好", start=0.0, end=1.0) w2 = SubtitleWord(text="世界", start=1.0, end=2.0) tl = SubtitleTimeline( segments=[ SubtitleSegment(text="你好", start=0.0, end=1.0, words=[w1]), SubtitleSegment(text="世界", start=1.0, end=2.0, words=[w2]), ] ) result = tl.merge_short_segments(min_chars=10) assert len(result.segments[0].words) == 2 assert result.segments[0].words[0].text == "你好" assert result.segments[0].words[1].text == "世界" def test_multiple_merged_groups(self): """多个合并组 — 短段会和后续段累积到够数才提交""" tl = SubtitleTimeline( segments=[ SubtitleSegment(text="一二三四五六七八", start=0, end=2), # 8字,够数,提交 SubtitleSegment(text="九", start=2, end=2.5), # 1字,入buffer SubtitleSegment(text="十", start=2.5, end=3), # 1字,入buffer(共2字) SubtitleSegment(text="一二三四五六七八九十", start=3, end=5), # 10字,入buffer后共12字,够数提交 ] ) result = tl.merge_short_segments(min_chars=8) # 第1段:"一二三四五六七八"(8字直接提交) # 第2段:"九十" + "一二三四五六七八九十" 累积到12字一起提交 assert result.segment_count == 2 assert result.segments[0].text == "一二三四五六七八" assert result.segments[1].text == "九十一二三四五六七八九十" def test_remaining_short_merged_with_last(self): """剩余短片段合并到最后一段""" tl = SubtitleTimeline( segments=[ SubtitleSegment(text="一二三四五六七八", start=0, end=2), # 8字 SubtitleSegment(text="一二三", start=2, end=3), # 3字,不够 ] ) result = tl.merge_short_segments(min_chars=8) # 最后的3字会合并到上一段(因为 < min_chars) assert result.segment_count == 1 assert result.segments[0].text == "一二三四五六七八一二三" def test_custom_min_chars(self): """自定义最小字数 — 累积到够数就提交,剩余短的合并到最后""" tl = SubtitleTimeline( segments=[ SubtitleSegment(text="一二", start=0, end=1), SubtitleSegment(text="三四", start=1, end=2), SubtitleSegment(text="五六", start=2, end=3), ] ) # min_chars=3: # "一二"(2字) → 不够 # +"三四"(共4字) → 够了,提交"一二三四",buffer清空 # "五六"(2字) → 循环结束,剩余 1 # 每段都不超过 max_chars(除了硬切的情况) for seg in result.segments: assert seg.char_count <= len(text) # 至少比原文短 def test_split_preserves_total_text(self): """拆分后总文本不变""" text = "你好世界。今天天气真好,我们出去玩吧!明天再见。" tl = SubtitleTimeline( segments=[ SubtitleSegment(text=text, start=0, end=10.0), ] ) result = tl.split_long_segments(max_chars=8) merged_text = "".join(s.text for s in result.segments) assert merged_text == text def test_split_time_proportional(self): """拆分后时间按字数比例分配""" text = "一二三四五六七八九十。" # 11字 tl = SubtitleTimeline( segments=[ SubtitleSegment(text=text, start=0, end=10.0), ] ) result = tl.split_long_segments(max_chars=5) # 总时长不变 assert result.segments[0].start == 0.0 assert result.segments[-1].end == pytest.approx(10.0) # 各段首尾相接 for i in range(len(result.segments) - 1): assert result.segments[i].end == pytest.approx(result.segments[i + 1].start) def test_split_with_words(self): """拆分时词级信息正确分配""" words = [ SubtitleWord(text="你好", start=0.0, end=1.0), SubtitleWord(text="世界", start=1.0, end=2.0), SubtitleWord(text="你好吗", start=2.0, end=3.5), ] tl = SubtitleTimeline( segments=[ SubtitleSegment(text="你好世界。你好吗?", start=0.0, end=3.5, words=words), ] ) result = tl.split_long_segments(max_chars=4) # 第一段应该有前几个词 assert len(result.segments) >= 2 total_words = sum(len(s.words) for s in result.segments) assert total_words == 3 # 词的总数不变 def test_multiple_mixed_segments(self): """混合长短片段""" tl = SubtitleTimeline( segments=[ SubtitleSegment(text="短", start=0, end=1), # 短 SubtitleSegment(text="一二三四五六七八九十一二三四五六七八九十", start=1, end=5), # 长 SubtitleSegment(text="也短", start=5, end=6), # 短 ] ) result = tl.split_long_segments(max_chars=10) assert result.segment_count >= 3 # 至少3段(中间被拆成多段) # 第一段还是原来的短的 assert result.segments[0].text == "短" # 最后一段还是原来的短的 assert result.segments[-1].text == "也短" def test_no_punctuation_hard_split(self): """没有标点时硬切""" text = "一二三四五六七八九十一二三四五六七八九十一二三四五" tl = SubtitleTimeline( segments=[ SubtitleSegment(text=text, start=0, end=10.0), ] ) result = tl.split_long_segments(max_chars=10) assert result.segment_count >= 3 for seg in result.segments: # 硬切的每段应该 <= max_chars assert seg.char_count <= 10 def test_preserves_language_and_duration(self): """拆分后保留语言和总时长""" tl = SubtitleTimeline( segments=[SubtitleSegment(text="a", start=0, end=1)], language="ja", total_duration=30.0, ) result = tl.split_long_segments(max_chars=20) assert result.language == "ja" assert result.total_duration == 30.0 def test_does_not_modify_original(self): """不修改原时间轴""" original_text = "一二三四五六七八九十一二三四五六七八九十" tl = SubtitleTimeline( segments=[ SubtitleSegment(text=original_text, start=0, end=5), ] ) result = tl.split_long_segments(max_chars=8) assert tl.segment_count == 1 assert tl.segments[0].text == original_text assert result is not tl class TestSplitTextByPunctuation: """_split_text_by_punctuation 静态方法测试""" def test_short_text_no_split(self): result = SubtitleTimeline._split_text_by_punctuation("你好世界", 10) assert result == ["你好世界"] def test_split_at_sentence_end(self): """在句末标点处断开""" result = SubtitleTimeline._split_text_by_punctuation("你好。世界。", 5) assert len(result) == 2 assert result[0] == "你好。" assert result[1] == "世界。" def test_split_at_comma(self): """在逗号处断开(超过最大长度时)""" text = "一二三四五六七八,二二三四五六七八。" result = SubtitleTimeline._split_text_by_punctuation(text, 10) assert len(result) >= 2 def test_no_punctuation_hard_split(self): """没有标点时硬切""" result = SubtitleTimeline._split_text_by_punctuation("一二三四五六七八九十", 5) assert len(result) == 2 assert result[0] == "一二三四五" assert result[1] == "六七八九十" def test_empty_text(self): # 空字符串循环不执行,current为空不append,返回空列表 result = SubtitleTimeline._split_text_by_punctuation("", 10) assert result == [] def test_mixed_punctuation(self): """混合标点""" text = "你好!吃饭了吗?是的,我吃过了。" result = SubtitleTimeline._split_text_by_punctuation(text, 6) # 验证所有段加起来等于原文 assert "".join(result) == text def test_sentence_end_with_min_length(self): """句末标点断句的「半长门槛」只在未超max_chars时生效; 超过max_chars回溯找标点时,即使首段很短也会断开。""" # "你好。" 3字 < max_chars//2(5),未超max_chars时不会主动断开 # 但加上后面的"世界很大很美好"后超过10字,回溯找标点找到"。",强制断开 text = "你好。世界很大很美好。" result = SubtitleTimeline._split_text_by_punctuation(text, 10) # 超过max_chars时回溯断开,首段可能很短 assert len(result) == 2 assert result[0] == "你好。" assert result[1] == "世界很大很美好。" # 总文本不变 assert "".join(result) == text def test_exclamation_and_question_marks(self): """感叹号和问号也算句末标点""" text = "你好吗!我很好!你呢?" result = SubtitleTimeline._split_text_by_punctuation(text, 4) assert len(result) >= 3 class TestMergeSegments: """_merge_segments 静态方法测试""" def test_merge_two_segments(self): result = SubtitleTimeline._merge_segments( [ SubtitleSegment(text="你好", start=0.0, end=1.0), SubtitleSegment(text="世界", start=1.0, end=2.0), ] ) assert result.text == "你好世界" assert result.start == 0.0 assert result.end == 2.0 def test_merge_empty_list(self): result = SubtitleTimeline._merge_segments([]) assert result.text == "" assert result.start == 0 assert result.end == 0 def test_merge_single_segment(self): seg = SubtitleSegment(text="test", start=1.0, end=2.0) result = SubtitleTimeline._merge_segments([seg]) assert result.text == "test" assert result.start == 1.0 assert result.end == 2.0 def test_merge_preserves_words(self): w1 = SubtitleWord(text="你好", start=0.0, end=1.0) w2 = SubtitleWord(text="世界", start=1.0, end=2.0) result = SubtitleTimeline._merge_segments( [ SubtitleSegment(text="你好", start=0.0, end=1.0, words=[w1]), SubtitleSegment(text="世界", start=1.0, end=2.0, words=[w2]), ] ) assert len(result.words) == 2 assert result.words[0].text == "你好" assert result.words[1].text == "世界" def test_merge_non_contiguous_segments(self): """合并非连续片段(有间隙)""" result = SubtitleTimeline._merge_segments( [ SubtitleSegment(text="a", start=0.0, end=1.0), SubtitleSegment(text="b", start=3.0, end=4.0), ] ) assert result.start == 0.0 assert result.end == 4.0 assert result.text == "ab" class TestMergeAndSplitRoundtrip: """合并和拆分的组合测试""" def test_split_then_merge_approximate(self): """拆分后再合并,总字数和总时长基本一致""" original_text = "你好世界。今天天气真好,我们出去玩吧!明天见。" tl = SubtitleTimeline( segments=[ SubtitleSegment(text=original_text, start=0.0, end=10.0), ] ) split = tl.split_long_segments(max_chars=5) merged = split.merge_short_segments(min_chars=50) # 足够大的min_chars让它们都合并 assert merged.segment_count == 1 assert merged.segments[0].text == original_text assert merged.segments[0].start == 0.0 assert merged.segments[0].end == pytest.approx(10.0)