File size: 7,462 Bytes
9010fc3
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
"""The server turn: one Japanese sentence in, one AvatarDirective out.

VOIC-03, VOIC-04 and VOIC-05 all land in ``blocks.turn``. These tests pin the three things that
are easy to get silently wrong: the directive's shape on the wire, the timeline agreeing with the
audio that is actually returned, and the "slower" re-read being a genuine re-synthesis rather than
a JS time-stretch (01-RESEARCH.md Pitfall 5).

Needs the VOICEVOX wheel and the LFS assets, like ``tests/test_tts_contract.py``. Skips cleanly
without them so a contributor's quick loop stays green.
"""

from __future__ import annotations

import ast
import base64
import json
from pathlib import Path

import pytest

pytest.importorskip("voicevox_core")

from japanese_avatar.telemetry.timings import TurnTimings  # noqa: E402
from japanese_avatar.ui.blocks import GREETING_TEXT, MAX_TEXT_CHARS, greeting, turn  # noqa: E402
from japanese_avatar.voice.models import AvatarDirective  # noqa: E402
from japanese_avatar.voice.tts import wav_duration_seconds  # noqa: E402
from japanese_avatar.voice.visemes import FRAMERATE, to_frame  # noqa: E402

SRC = Path(__file__).resolve().parents[1] / "src" / "japanese_avatar"

SHORT_TEXT = "こんにちは"
DIRECTIVE_KEYS = {"turn_id", "audio_url", "timeline", "subtitle", "expression", "speed", "timings"}
STAGE_KEYS = {"audio_query_ms", "synthesis_ms", "timeline_ms", "encode_ms", "server_total_ms"}
ONE_FRAME = 1 / FRAMERATE
DATA_URL_PREFIX = "data:audio/wav;base64,"


def _decode_wav(directive: dict) -> bytes:
    assert directive["audio_url"].startswith(DATA_URL_PREFIX)
    return base64.b64decode(directive["audio_url"][len(DATA_URL_PREFIX) :])


def _total(directive: dict) -> float:
    return sum(event["dur"] for event in directive["timeline"])


@pytest.fixture(scope="module")
def normal() -> dict:
    return turn(SHORT_TEXT)


@pytest.fixture(scope="module")
def slow() -> dict:
    return turn(SHORT_TEXT, speed=0.75)


def test_directive_json_roundtrips(normal):
    directive = AvatarDirective(**normal)
    parsed = json.loads(directive.to_json())
    assert set(parsed) == DIRECTIVE_KEYS
    assert set(normal) == DIRECTIVE_KEYS
    assert parsed["subtitle"] == SHORT_TEXT
    assert parsed["expression"] == "neutral"
    assert parsed["speed"] == 1.0
    assert parsed["turn_id"].startswith("t-") and len(parsed["turn_id"]) == 10


def test_directive_audio_is_data_url(normal):
    wav = _decode_wav(normal)
    assert wav[:4] == b"RIFF"
    assert len(wav) > 1000


def test_turn_returns_matching_timeline_duration(normal):
    """The timeline must end where the audio ends, within one VOICEVOX frame."""
    duration = wav_duration_seconds(_decode_wav(normal))
    assert abs(_total(normal) - duration) <= ONE_FRAME, (
        f"timeline totals {_total(normal):.6f}s but the returned WAV is {duration:.6f}s"
    )
    for event in normal["timeline"]:
        assert set(event) == {"t", "dur", "viseme", "weight"}


def test_slower_turn_rebuilds_timeline(normal, slow):
    """VOIC-03 and Pitfall 5: the slow timeline is built from the slow query.

    (a) same viseme sequence, (b) the slow total agrees with the SLOW WAV to one frame and its
    ratio to the normal total is within 2% of 1/0.75 - plan 01-06 measured that VOICEVOX
    re-quantises after dividing, so the realised ratio lands near 1.3333, never exactly on it -
    and (c) the first event is the pre-phoneme silence quantised at speedScale 0.75, which a
    timeline scaled in JS from the fast one could not reproduce.
    """
    assert slow["speed"] == 0.75
    assert normal["speed"] == 1.0

    normal_seq = [e["viseme"] for e in normal["timeline"]]
    slow_seq = [e["viseme"] for e in slow["timeline"]]
    assert normal_seq == slow_seq

    slow_duration = wav_duration_seconds(_decode_wav(slow))
    assert abs(_total(slow) - slow_duration) <= ONE_FRAME, (
        f"slow timeline totals {_total(slow):.6f}s but the slow WAV is {slow_duration:.6f}s"
    )
    ratio = _total(slow) / _total(normal)
    assert abs(ratio - 1 / 0.75) < 0.02 * (1 / 0.75), ratio

    # 0.1 s of pre-phoneme silence is VOICEVOX's default; a rebuilt query carries it through
    # frames_for() at 0.75: round(round(0.1 * 93.75) / 0.75) = round(9 / 0.75) = 12 frames.
    first_normal = normal["timeline"][0]["dur"]
    first_slow = slow["timeline"][0]["dur"]
    pre_frames = to_frame(first_normal)
    assert first_slow == pytest.approx(round(pre_frames / 0.75) / FRAMERATE, abs=1e-9)
    assert first_slow > first_normal


def test_timings_are_per_request(normal, slow):
    """Two turns, two independent timing records - and no module-level TurnTimings anywhere."""
    assert normal["timings"] is not slow["timings"]
    assert set(normal["timings"]) >= STAGE_KEYS
    assert set(slow["timings"]) >= STAGE_KEYS
    for key in STAGE_KEYS:
        assert normal["timings"][key] >= 0
        assert slow["timings"][key] >= 0
    assert normal["timings"]["server_total_ms"] > 0
    # Distinct objects with independent values: mutating one must not touch the other.
    normal["timings"]["probe"] = 1
    assert "probe" not in slow["timings"]
    del normal["timings"]["probe"]

    for path in sorted(SRC.rglob("*.py")):
        tree = ast.parse(path.read_text(encoding="utf-8"), filename=str(path))
        for node in tree.body:
            targets = []
            if isinstance(node, ast.Assign | ast.AnnAssign) and node.value is not None:
                targets = [node.value]
            for value in targets:
                if isinstance(value, ast.Call):
                    name = value.func
                    called = name.id if isinstance(name, ast.Name) else getattr(name, "attr", "")
                    assert called != "TurnTimings", (
                        f"{path.relative_to(SRC.parent)} holds a module-level TurnTimings; Gradio "
                        "shares module globals across every visitor session"
                    )


def test_turn_rejects_bad_input_with_a_structured_error():
    """The bridge turns a raised exception into `undefined` in the browser, so never raise."""
    assert "error" in turn("")
    assert "error" in turn("   ")
    assert "error" in turn("あ" * (MAX_TEXT_CHARS + 1))
    assert "error" in turn(SHORT_TEXT, speed=0.0)
    assert "error" in turn(None)


def test_turn_accepts_the_bridge_payload_shapes():
    """gr.HTML's server bridge passes ONE JSON argument: a dict from JS, or a list for multi-arg."""
    from_dict = turn({"text": SHORT_TEXT, "speed": 0.75})
    assert from_dict["speed"] == 0.75 and from_dict["subtitle"] == SHORT_TEXT
    from_list = turn([SHORT_TEXT, 0.75])
    assert from_list["speed"] == 0.75


def test_greeting_speaks_with_no_input():
    """`server.greeting()` from JS arrives as greeting([]) - the positional must be tolerated."""
    directive = greeting([])
    assert set(directive) == DIRECTIVE_KEYS
    assert directive["subtitle"] == GREETING_TEXT
    assert greeting()["subtitle"] == GREETING_TEXT


def test_turn_timings_stage_and_mark_share_one_record():
    timings = TurnTimings()
    timings.mark("audio_query", 1.5)
    timings.mark("audio_query", 0.5)
    with timings.stage("encode"):
        pass
    out = timings.as_dict()
    assert out["audio_query_ms"] == 2.0
    assert out["encode_ms"] >= 0
    assert out["server_total_ms"] >= out["encode_ms"]
    out["audio_query_ms"] = 99
    assert timings.as_dict()["audio_query_ms"] == 2.0, "as_dict must return a copy"