Spaces:
Sleeping
Sleeping
File size: 3,685 Bytes
2e818da | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 | 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"
|