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"