File size: 2,392 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
import pytest

from app.services.chat_tts_stream import ChatTTSStream


class RecordingManager:
    def __init__(self) -> None:
        self.begin_calls: list[dict] = []
        self.enqueued: list[tuple[str, int, str]] = []
        self.finished: list[str] = []

    async def begin(self, project_id, message_id, send, **kwargs):
        self.begin_calls.append(kwargs)
        return "generation"

    async def enqueue(self, generation_id, sentence_index, text):
        self.enqueued.append((generation_id, sentence_index, text))

    async def finish(self, generation_id):
        self.finished.append(generation_id)


@pytest.mark.asyncio
async def test_voice_originated_stream_buffers_one_complete_native_utterance():
    manager = RecordingManager()
    stream = await ChatTTSStream.create(
        manager=manager,
        project_id="project",
        message_id="message",
        send=lambda *_args: None,
        input_mode="voice",
        auto_read_reply=True,
    )

    await stream.feed("First sentence. Sec")
    await stream.feed("ond sentence!")
    assert manager.enqueued == []
    await stream.finish()

    assert manager.enqueued == [("generation", 0, "First sentence. Second sentence!")]
    assert manager.finished == ["generation"]


@pytest.mark.asyncio
async def test_auto_read_forwards_the_model_and_voice_selection():
    manager = RecordingManager()

    await ChatTTSStream.create(
        manager=manager,
        project_id="project",
        message_id="message",
        send=lambda *_args: None,
        input_mode="voice",
        auto_read_reply=True,
        model_id="piper",
        voice_id="en_US-lessac-medium",
    )

    assert manager.begin_calls == [{"model_id": "piper", "voice_id": "en_US-lessac-medium"}]


@pytest.mark.asyncio
@pytest.mark.parametrize(("input_mode", "auto_read_reply"), [("text", True), ("voice", False)])
async def test_stream_is_disabled_for_text_or_opted_out_turns(input_mode, auto_read_reply):
    manager = RecordingManager()
    stream = await ChatTTSStream.create(
        manager=manager,
        project_id="project",
        message_id="message",
        send=lambda *_args: None,
        input_mode=input_mode,
        auto_read_reply=auto_read_reply,
    )
    await stream.feed("Nothing should play.")
    await stream.finish()

    assert stream.generation_id is None
    assert manager.begin_calls == []