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 @pytest.mark.asyncio 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"