study-buddy / tests /test_tts_service.py
GitHub Actions
deploy d092bea3608b7a29952f16357fda39b7a29e399b
2e818da
Raw
History Blame Contribute Delete
3.69 kB
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"