Spaces:
Sleeping
Sleeping
| import pytest | |
| from app.services.tts_service import ( | |
| ModelManager, | |
| ProviderDescriptor, | |
| SynthesizerRegistry, | |
| TTSGenerationManager, | |
| VoiceDescriptor, | |
| build_default_registry, | |
| ) | |
| class FakeSynthesizer: | |
| sample_rate = 24_000 | |
| channels = 1 | |
| def __init__(self, model_id: str, voice_id: str) -> None: | |
| self.model_id = model_id | |
| self.voice_id = voice_id | |
| self.unloaded = False | |
| def stream_pcm(self, text: str): | |
| yield f"{text}:1".encode() | |
| yield f"{text}:2".encode() | |
| def unload(self) -> None: | |
| self.unloaded = True | |
| def _test_registry(created: list) -> SynthesizerRegistry: | |
| registry = SynthesizerRegistry() | |
| for model_id in ("model-a", "model-b"): | |
| registry.register( | |
| ProviderDescriptor( | |
| model_id=model_id, | |
| display_name=model_id, | |
| license="MIT", | |
| sample_rate=24_000, | |
| channels=1, | |
| voices=(VoiceDescriptor("voice", "Voice"),), | |
| ), | |
| lambda voice_id, model_id=model_id: created.append(FakeSynthesizer(model_id, voice_id)) or created[-1], | |
| ) | |
| return registry | |
| def test_default_registry_exposes_only_pocket_tts_and_piper(): | |
| catalog = build_default_registry().catalog() | |
| assert {model["model_id"] for model in catalog} == {"pocket-tts", "piper"} | |
| def test_pocket_tts_default_voice_is_the_bundled_preset(): | |
| catalog = build_default_registry().catalog() | |
| pocket = next(model for model in catalog if model["model_id"] == "pocket-tts") | |
| assert [voice["voice_id"] for voice in pocket["voices"]] == ["anshuman-normal-custom"] | |
| def test_piper_exposes_exactly_one_bundled_voice(): | |
| catalog = build_default_registry().catalog() | |
| piper = next(model for model in catalog if model["model_id"] == "piper") | |
| assert [voice["voice_id"] for voice in piper["voices"]] == ["en_US-lessac-medium"] | |
| def test_registry_create_rejects_unknown_voice(): | |
| registry = SynthesizerRegistry() | |
| registry.register( | |
| ProviderDescriptor("m", "M", "MIT", 24_000, 1, (VoiceDescriptor("v", "V"),)), | |
| lambda voice: FakeSynthesizer("m", voice), | |
| ) | |
| with pytest.raises(ValueError): | |
| registry.create("m", "not-a-voice") | |
| def test_model_manager_reuses_warm_synthesizer_for_same_selection(): | |
| created: list = [] | |
| manager = ModelManager(_test_registry(created)) | |
| first = manager.select("model-a", "voice") | |
| second = manager.select("model-a", "voice") | |
| assert first is second | |
| assert manager.last_selection_was_reused is True | |
| def test_model_manager_unloads_previous_synthesizer_on_switch(): | |
| created: list = [] | |
| manager = ModelManager(_test_registry(created)) | |
| first = manager.select("model-a", "voice") | |
| manager.select("model-b", "voice") | |
| assert first.unloaded is True | |
| async def test_generation_manager_streams_chunks_then_done(): | |
| created: list = [] | |
| manager = TTSGenerationManager(model_manager=ModelManager(_test_registry(created))) | |
| events: list[tuple[str, dict]] = [] | |
| def send(event, payload): | |
| events.append((event, payload)) | |
| generation_id = await manager.begin("proj", "msg", send, model_id="model-a", voice_id="voice") | |
| await manager.enqueue(generation_id, 0, "hello") | |
| await manager.finish(generation_id) | |
| await manager.wait(generation_id) | |
| event_types = [event for event, _ in events] | |
| assert event_types[0] == "TTS_START" | |
| assert "TTS_CHUNK" in event_types | |
| assert event_types[-1] == "TTS_DONE" | |
| done_payload = next(payload for event, payload in events if event == "TTS_DONE") | |
| assert done_payload["status"] == "success" | |