diff --git a/tests/unit/test_mock_asr_service.py b/tests/unit/test_mock_asr_service.py new file mode 100755 index 000000000..6491c310d --- /dev/null +++ b/tests/unit/test_mock_asr_service.py @@ -0,0 +1,191 @@ +"""MockASRService 单测 — Mock ASR服务纯逻辑部分.""" +from __future__ import annotations + +from pathlib import Path + +import pytest + +from packages.adapters.asr.mock_asr_service import MockASRService +from packages.ports.asr_service import ASRServiceError + + +# ── Fixtures ──────────────────────────────────────────────────────────────── + + +@pytest.fixture +def service(): + return MockASRService() + + +# ── transcribe 基础行为 ──────────────────────────────────────────────────────── + + +class TestTranscribeBasic: + """transcribe 基本行为.""" + + def test_audio_not_found_raises(self, service, tmp_path): + """音频文件不存在时抛 ASRServiceError.""" + with pytest.raises(ASRServiceError): + service.transcribe(tmp_path / "nonexistent.wav") + + def test_default_mock_text_produces_timeline(self, service, tmp_path): + """不传 mock_text 时生成默认测试字幕.""" + audio = tmp_path / "test.wav" + audio.write_bytes(b"fake audio") + + timeline = service.transcribe(audio) + + assert timeline is not None + assert timeline.language == "zh" + assert timeline.total_duration > 0 + assert len(timeline.segments) > 0 + + def test_custom_mock_text(self, tmp_path): + """使用自定义 mock_text.""" + text = "你好世界。今天天气真好!" + service = MockASRService(mock_text=text) + audio = tmp_path / "test.wav" + audio.write_bytes(b"fake audio") + + timeline = service.transcribe(audio) + + assert timeline.total_duration > 0 + # 两句话,应该有2个segment + assert len(timeline.segments) == 2 + assert "你好世界" in timeline.segments[0].text + assert "今天天气真好" in timeline.segments[1].text + + def test_reads_txt_file(self, tmp_path): + """音频同目录同名 txt 文件存在时,读取其内容.""" + audio = tmp_path / "voice.wav" + audio.write_bytes(b"fake audio") + txt = tmp_path / "voice.txt" + txt.write_text("这是从文件读取的。内容。", encoding="utf-8") + + service = MockASRService() + timeline = service.transcribe(audio) + + assert len(timeline.segments) == 2 + assert "从文件读取" in timeline.segments[0].text + + def test_mock_text_takes_priority_over_txt(self, tmp_path): + """mock_text 参数优先于 txt 文件.""" + audio = tmp_path / "test.wav" + audio.write_bytes(b"fake audio") + txt = tmp_path / "test.txt" + txt.write_text("文件里的文字。", encoding="utf-8") + + service = MockASRService(mock_text="自定义的。优先!") + timeline = service.transcribe(audio) + + assert "自定义的" in timeline.segments[0].text + assert "文件" not in timeline.segments[0].text + + def test_custom_language(self, service, tmp_path): + """可以指定语言.""" + audio = tmp_path / "test.wav" + audio.write_bytes(b"fake audio") + timeline = service.transcribe(audio, language="en") + assert timeline.language == "en" + + +# ── _text_to_segments 切分逻辑 ────────────────────────────────────────── + + +class TestTextToSegments: + """_text_to_segments 文本切分逻辑.""" + + def setup_method(self): + self.service = MockASRService() + + def test_split_by_chinese_period(self): + """按中文句号切分.""" + text = "第一句。第二句。第三句。" + segments = self.service._text_to_segments(text, 10.0, with_word_timestamps=False) + assert len(segments) == 3 + + def test_split_by_exclamation(self): + """按感叹号切分.""" + text = "你好!世界!" + segments = self.service._text_to_segments(text, 5.0, with_word_timestamps=False) + assert len(segments) == 2 + + def test_split_by_question_mark(self): + """按问号切分.""" + text = "你好吗?我很好。" + segments = self.service._text_to_segments(text, 5.0, with_word_timestamps=False) + assert len(segments) == 2 + + def test_mixed_punctuation(self): + """混合标点.""" + text = "你好!你是谁?我是测试。再见!" + segments = self.service._text_to_segments(text, 10.0, with_word_timestamps=False) + assert len(segments) == 4 + + def test_empty_text_returns_empty(self): + """空文本返回空列表.""" + segments = self.service._text_to_segments("", 5.0, with_word_timestamps=False) + assert segments == [] + + def test_no_punctuation_single_segment(self): + """没有标点时整段作为一个segment.""" + text = "这是一段没有标点的文字" + segments = self.service._text_to_segments(text, 5.0, with_word_timestamps=False) + assert len(segments) == 1 + assert segments[0].text == text + + def test_segments_duration_adds_up(self): + """所有segment时长加起来约等于总时长.""" + text = "第一句。第二句。" + total = 10.0 + segments = self.service._text_to_segments(text, total, with_word_timestamps=False) + sum_duration = sum(s.end - s.start for s in segments) + assert abs(sum_duration - total) < 0.01 + + def test_segments_sequential(self): + """segment按顺序排列,首尾相接.""" + text = "第一句。第二句。第三句。" + segments = self.service._text_to_segments(text, 9.0, with_word_timestamps=False) + + assert segments[0].start == 0.0 + for i in range(1, len(segments)): + assert abs(segments[i].start - segments[i - 1].end) < 0.001 + + def test_with_word_timestamps(self): + """带词级时间戳时,每个字一个word.""" + text = "你好世界!" + segments = self.service._text_to_segments(text, 2.0, with_word_timestamps=True) + + assert len(segments) == 1 + # 4个汉字 + 1个感叹号 = 5个字符 + assert len(segments[0].words) == 5 + + def test_word_timestamps_sequential(self): + """词级时间戳按顺序排列.""" + text = "你好!" + segments = self.service._text_to_segments(text, 3.0, with_word_timestamps=True) + + words = segments[0].words + assert len(words) == 3 + assert words[0].start == 0.0 + # 最后一个词的结束时间约等于 segment 结束时间 + assert abs(words[-1].end - segments[0].end) < 0.01 + + def test_without_word_timestamps(self): + """不带词级时间戳时,words为空.""" + text = "你好世界。" + segments = self.service._text_to_segments(text, 2.0, with_word_timestamps=False) + + assert len(segments) == 1 + assert segments[0].words == [] + + def test_duration_proportional_to_length(self): + """长句子占时长,短句子占时短.""" + text = "短。很长很长很长很长的句子。" + segments = self.service._text_to_segments(text, 10.0, with_word_timestamps=False) + + # 第一句1个字,第二句9个字(包括标点) + # 第一句时长应该比第二句短 + dur0 = segments[0].end - segments[0].start + dur1 = segments[1].end - segments[1].start + assert dur1 > dur0 # 第二句更长 diff --git a/tests/unit/test_mock_tts_service.py b/tests/unit/test_mock_tts_service.py new file mode 100755 index 000000000..59c9cfd9f --- /dev/null +++ b/tests/unit/test_mock_tts_service.py @@ -0,0 +1,143 @@ +"""MockTtsService 单测 — Mock TTS服务纯逻辑部分.""" +from __future__ import annotations + +import pytest + +from packages.adapters.tts.mock_tts_service import MockTtsService, _CHARS_PER_SECOND + + +# ── Fixtures ──────────────────────────────────────────────────────────────── + + +@pytest.fixture +def service(): + return MockTtsService() + + +# ── estimate_duration ────────────────────────────────────────────────────── + + +class TestEstimateDuration: + """estimate_duration 时长估算.""" + + def test_empty_text_returns_zero(self, service): + assert service.estimate_duration("") == 0.0 + + def test_whitespace_only_returns_zero(self, service): + assert service.estimate_duration(" \n\t ") == 0.0 + + def test_single_char(self, service): + result = service.estimate_duration("你") + assert abs(result - 1.0 / _CHARS_PER_SECOND) < 0.001 + + def test_default_speed(self, service): + """默认 speed=1.0.""" + text = "你好世界" # 4个字 + result = service.estimate_duration(text) + expected = 4.0 / _CHARS_PER_SECOND + assert abs(result - expected) < 0.001 + + def test_faster_speed_shortens_duration(self, service): + """语速越快,时长越短.""" + text = "你好世界" + normal = service.estimate_duration(text, speed=1.0) + fast = service.estimate_duration(text, speed=2.0) + assert fast < normal + assert abs(fast - normal / 2) < 0.001 + + def test_slower_speed_lengthens_duration(self, service): + """语速越慢,时长越长.""" + text = "你好世界" + normal = service.estimate_duration(text, speed=1.0) + slow = service.estimate_duration(text, speed=0.5) + assert slow > normal + assert abs(slow - normal / 0.5) < 0.001 + + def test_speed_clamped_at_minimum(self, service): + """speed < 0.1 时被钳制到 0.1,避免除零.""" + text = "你好" + # 传一个极小的值,不应该崩溃,且时长不会无限大 + result = service.estimate_duration(text, speed=0.001) + assert result > 0 + # 应该等同于 speed=0.1 + expected = 2.0 / _CHARS_PER_SECOND / 0.1 + assert abs(result - expected) < 0.001 + + def test_chinese_and_english_mixed(self, service): + """中英文混合时按非空白字符计数.""" + text = "Hello 世界" # H-e-l-l-o + 世-界 = 7个非空白字符 + result = service.estimate_duration(text) + expected = 7.0 / _CHARS_PER_SECOND + assert abs(result - expected) < 0.001 + + def test_negative_speed(self, service): + """负语速按最小处理(取 max(0.1, speed)).""" + text = "你好" + result = service.estimate_duration(text, speed=-2.0) + assert result > 0 + # 等同于 speed=0.1 + expected = 2.0 / _CHARS_PER_SECOND / 0.1 + assert abs(result - expected) < 0.001 + + +# ── _extract_freq ────────────────────────────────────────────────────────── + + +class TestExtractFreq: + """_extract_freq 基频提取.""" + + def test_sine_prefix_returns_freq(self, service): + """sine_ 前缀的voice_id,从下划线后提取频率.""" + result = service._extract_freq("sine_440", "female") + assert result == 440.0 + + def test_sine_with_decimal(self, service): + """支持小数频率.""" + result = service._extract_freq("sine_261.63", "female") + assert abs(result - 261.63) < 0.001 + + def test_sine_invalid_number_falls_back(self, service): + """sine_ 后面不是数字时,fallback 到性别默认值.""" + result = service._extract_freq("sine_abc", "female") + assert result == 220.0 # female 默认 + + def test_sine_no_number_falls_back(self, service): + """sine_ 后面没有内容时,fallback.""" + result = service._extract_freq("sine_", "male") + assert result == 120.0 # male 默认 + + def test_male_default(self, service): + result = service._extract_freq("some_voice", "male") + assert result == 120.0 + + def test_female_default(self, service): + result = service._extract_freq("some_voice", "female") + assert result == 220.0 + + def test_child_default(self, service): + result = service._extract_freq("some_voice", "child") + assert result == 350.0 + + def test_unknown_gender_defaults_to_female(self, service): + """未知性别 fallback 到 female.""" + result = service._extract_freq("some_voice", "alien") + assert result == 220.0 + + def test_empty_gender_defaults_to_female(self, service): + result = service._extract_freq("some_voice", "") + assert result == 220.0 + + +# ── provider_name / available_voices ─────────────────────────────────────── + + +class TestProviderInfo: + """provider_name 和 available_voices.""" + + def test_provider_name(self, service): + assert service.provider_name == "mock" + + def test_available_voices_returns_list(self, service): + voices = service.available_voices() + assert isinstance(voices, list) + assert len(voices) > 0 diff --git a/tests/unit/test_tts_service_factory.py b/tests/unit/test_tts_service_factory.py new file mode 100755 index 000000000..22f29f0cd --- /dev/null +++ b/tests/unit/test_tts_service_factory.py @@ -0,0 +1,138 @@ +"""TTS Service Factory 单测 — TTS服务工厂.""" +from __future__ import annotations + +import os +from unittest.mock import patch + +import pytest + +# 注意:工厂模块有全局状态(_PROVIDERS),每个测试前重置 +from services.tts_service_factory import ( + _PROVIDERS, + available_providers, + get_tts_service, + register_provider, +) + + +# ── Fixtures ──────────────────────────────────────────────────────────────── + + +@pytest.fixture(autouse=True) +def reset_providers(): + """每个测试前后重置 provider 注册表.""" + # 保存原始状态 + original = dict(_PROVIDERS) + yield + # 恢复 + _PROVIDERS.clear() + _PROVIDERS.update(original) + + +# ── register_provider ────────────────────────────────────────────────────── + + +class TestRegisterProvider: + """register_provider 注册供应商.""" + + def test_register_new_provider(self): + class DummyService: + pass + + register_provider("dummy", DummyService) + assert "dummy" in _PROVIDERS + assert _PROVIDERS["dummy"] is DummyService + + def test_register_overwrites_existing(self): + class ServiceV1: + pass + + class ServiceV2: + pass + + register_provider("test", ServiceV1) + register_provider("test", ServiceV2) + assert _PROVIDERS["test"] is ServiceV2 + + +# ── get_tts_service ─────────────────────────────────────────────────────── + + +class TestGetTtsService: + """get_tts_service 获取TTS服务.""" + + def test_get_mock_provider(self): + """mock 供应商可用.""" + service = get_tts_service("mock") + assert service is not None + assert service.provider_name == "mock" + + def test_get_cosyvoice_provider(self): + """cosyvoice 别名映射到 CosyVoiceTtsService.""" + # 不传 api_key 也能实例化(默认空字符串) + service = get_tts_service("cosyvoice") + assert service is not None + assert service.provider_name == "cosyvoice" + + def test_aliyun_alias_maps_to_cosyvoice(self): + """aliyun 是 cosyvoice 的别名.""" + service = get_tts_service("aliyun") + assert service.provider_name == "cosyvoice" + + def test_dashscope_alias_maps_to_cosyvoice(self): + """dashscope 是 cosyvoice 的别名.""" + service = get_tts_service("dashscope") + assert service.provider_name == "cosyvoice" + + def test_unknown_provider_falls_back_to_mock(self): + """未知供应商回退到 mock.""" + service = get_tts_service("unknown_provider_xyz") + assert service.provider_name == "mock" + + def test_provider_name_case_insensitive(self): + """供应商名称不区分大小写.""" + service = get_tts_service("MOCK") + assert service.provider_name == "mock" + + def test_passes_kwargs_to_constructor(self): + """kwargs 传递给服务构造函数.""" + # MockTtsService 接受 ffmpeg_bin 参数 + service = get_tts_service("mock", ffmpeg_bin="/custom/ffmpeg") + assert service is not None + + def test_none_provider_reads_env_var(self): + """provider=None 时从 TTS_PROVIDER 环境变量读取.""" + with patch.dict(os.environ, {"TTS_PROVIDER": "mock"}): + service = get_tts_service(None) + assert service.provider_name == "mock" + + def test_empty_env_falls_back_to_auto_detect(self): + """环境变量为空时自动检测.""" + with patch.dict(os.environ, {"TTS_PROVIDER": ""}): + # 没有 cosyvoice_api_key 时应该用 mock + service = get_tts_service(None) + assert service.provider_name == "mock" + + +# ── available_providers ──────────────────────────────────────────────────── + + +class TestAvailableProviders: + """available_providers 可用供应商列表.""" + + def test_returns_list(self): + result = available_providers() + assert isinstance(result, list) + assert len(result) >= 1 # 至少有 mock + + def test_mock_is_always_available(self): + result = available_providers() + assert "mock" in result + + def test_after_register_appears_in_list(self): + class Dummy: + pass + + register_provider("dummy_test", Dummy) + result = available_providers() + assert "dummy_test" in result