test(P3-1): 第61波 mock ASR/TTS + TTS工厂单测(+52) #854
Executable
+191
@@ -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 # 第二句更长
|
||||
Executable
+143
@@ -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
|
||||
Executable
+138
@@ -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
|
||||
Reference in New Issue
Block a user