408 lines
16 KiB
Python
408 lines
16 KiB
Python
from __future__ import annotations
|
|
|
|
import asyncio
|
|
import base64
|
|
import json
|
|
import os
|
|
import unittest
|
|
from types import SimpleNamespace
|
|
from unittest import mock
|
|
|
|
import aiohttp
|
|
|
|
from app.livekit.adapters import xai_tts as xai_tts_module
|
|
from app.livekit.adapters.xai_tts import (
|
|
AUTH_METHOD_API_KEY,
|
|
DEFAULT_LANGUAGE,
|
|
DEFAULT_VOICE,
|
|
OraclexAITTS,
|
|
)
|
|
|
|
|
|
def _event(event_type: str, **values: object) -> SimpleNamespace:
|
|
return SimpleNamespace(
|
|
type=aiohttp.WSMsgType.TEXT,
|
|
data=json.dumps({"type": event_type, **values}),
|
|
)
|
|
|
|
|
|
class _FakeWebSocket:
|
|
def __init__(self, events: list[SimpleNamespace]) -> None:
|
|
self.events = list(events)
|
|
self.sent: list[dict[str, object]] = []
|
|
self.closed = False
|
|
|
|
def exception(self):
|
|
return None
|
|
|
|
async def send_str(self, payload: str) -> None:
|
|
self.sent.append(json.loads(payload))
|
|
|
|
async def receive(self) -> SimpleNamespace:
|
|
if not self.events:
|
|
await asyncio.sleep(60)
|
|
return self.events.pop(0)
|
|
|
|
async def close(self) -> None:
|
|
self.closed = True
|
|
|
|
|
|
class _Emitter:
|
|
def __init__(self) -> None:
|
|
self.audio = bytearray()
|
|
|
|
def push(self, payload: bytes) -> None:
|
|
self.audio.extend(payload)
|
|
|
|
|
|
class _Stream:
|
|
_segment_id = "segment-test"
|
|
|
|
def _mark_started(self) -> None:
|
|
return None
|
|
|
|
def _note_provider_ttfb(self, _provider_ttfb: float) -> None:
|
|
return None
|
|
|
|
|
|
class _EventOwner:
|
|
def __init__(self) -> None:
|
|
self.events: list[tuple[str, dict[str, object]]] = []
|
|
|
|
def _take_initial_greeting_capture(self, _text: str) -> None:
|
|
return None
|
|
|
|
def emit(self, event_name: str, event: dict[str, object]) -> None:
|
|
self.events.append((event_name, event))
|
|
|
|
|
|
class XAITTSUpgradeTests(unittest.IsolatedAsyncioTestCase):
|
|
def test_underflow_limit_defaults_to_one_second(self) -> None:
|
|
with mock.patch.dict(os.environ, {}, clear=True):
|
|
self.assertEqual(xai_tts_module._underflow_error_ms(), 1000)
|
|
self.assertEqual(xai_tts_module._turn_total_timeout_s(), 60.0)
|
|
|
|
def test_estimated_pcm_balance_uses_first_pcm_release_without_prebuffer(self) -> None:
|
|
self.assertEqual(
|
|
xai_tts_module._estimated_pcm_balance_s(
|
|
pcm_duration_s=1.0, first_pcm_released_at=100.0, now=101.125
|
|
),
|
|
-0.125,
|
|
)
|
|
|
|
def _connection(
|
|
self, events: list[SimpleNamespace], owner: _EventOwner | None = None
|
|
):
|
|
options = xai_tts_module._TTSOptions(
|
|
base_url="wss://example.test/tts",
|
|
voice=DEFAULT_VOICE,
|
|
language=DEFAULT_LANGUAGE,
|
|
)
|
|
auth = xai_tts_module._AuthOptions(
|
|
method=AUTH_METHOD_API_KEY,
|
|
api_key="key-123",
|
|
)
|
|
connection = xai_tts_module._Connection(
|
|
opts=options,
|
|
auth=auth,
|
|
session=object(),
|
|
owner=owner,
|
|
)
|
|
connection._ws = _FakeWebSocket(events)
|
|
connection._note_activity()
|
|
return connection
|
|
|
|
async def _synthesize(self, events: list[SimpleNamespace]):
|
|
connection = self._connection(events)
|
|
emitter = _Emitter()
|
|
result = await connection.synthesize_turn(
|
|
"nova fala",
|
|
output_emitter=emitter,
|
|
stream=_Stream(),
|
|
timeout=1.0,
|
|
turn_index=0,
|
|
connection_reused=True,
|
|
)
|
|
return connection, emitter, result
|
|
|
|
def test_legacy_public_name_and_websocket_url_are_preserved(self) -> None:
|
|
tts = OraclexAITTS(api_key="key-123", websocket_url="wss://legacy.test/tts")
|
|
|
|
self.assertIs(xai_tts_module.TTS, OraclexAITTS)
|
|
self.assertEqual(tts._opts.base_url, "wss://legacy.test/tts")
|
|
self.assertEqual(tts._opts.voice, DEFAULT_VOICE)
|
|
self.assertEqual(tts._opts.language, DEFAULT_LANGUAGE)
|
|
self.assertEqual(tts.model, DEFAULT_VOICE)
|
|
|
|
def test_api_key_auth_and_iam_auth_validation_are_available(self) -> None:
|
|
with mock.patch.dict(os.environ, {"XAI_API_KEY": "env-key"}, clear=True):
|
|
tts = OraclexAITTS()
|
|
|
|
self.assertEqual(tts._auth.method, AUTH_METHOD_API_KEY)
|
|
self.assertEqual(tts._auth.api_key, "env-key")
|
|
self.assertEqual(
|
|
xai_tts_module._request_headers(tts._auth, "wss://example.test/tts"),
|
|
{"Authorization": "Bearer env-key"},
|
|
)
|
|
|
|
with self.assertRaisesRegex(ValueError, "compartment_id"):
|
|
OraclexAITTS(auth_method="INSTANCE_PRINCIPAL")
|
|
|
|
def test_cached_greeting_read_failure_falls_back_to_tts(self) -> None:
|
|
class BrokenCache:
|
|
def __init__(self) -> None:
|
|
self.key = SimpleNamespace(digest="cache-key")
|
|
self.discarded = []
|
|
|
|
def key_for(self, **_kwargs):
|
|
return self.key
|
|
|
|
def has(self, _key) -> bool:
|
|
return True
|
|
|
|
def frames(self, _key):
|
|
raise FileNotFoundError("cached WAV disappeared")
|
|
|
|
def discard(self, key) -> None:
|
|
self.discarded.append(key)
|
|
|
|
cache = BrokenCache()
|
|
tts = OraclexAITTS(
|
|
api_key="key-123",
|
|
websocket_url="wss://example.test/tts",
|
|
initial_greeting_audio_cache=cache,
|
|
initial_greeting_agent="conta",
|
|
)
|
|
|
|
self.assertIsNone(tts.initial_greeting_audio("Olá, como posso ajudar?"))
|
|
self.assertEqual(cache.discarded, [cache.key])
|
|
self.assertIs(tts._initial_greeting_capture_key, cache.key)
|
|
|
|
def test_cached_greeting_hit_is_logged(self) -> None:
|
|
cache = mock.Mock()
|
|
cache_key = SimpleNamespace(digest="cache-key")
|
|
cached_audio = object()
|
|
cache.key_for.return_value = cache_key
|
|
cache.has.return_value = True
|
|
cache.frames.return_value = cached_audio
|
|
tts = OraclexAITTS(
|
|
api_key="key-123",
|
|
websocket_url="wss://example.test/tts",
|
|
initial_greeting_audio_cache=cache,
|
|
initial_greeting_agent="conta",
|
|
)
|
|
|
|
with mock.patch.object(xai_tts_module, "_runtime_logger") as runtime_logger:
|
|
self.assertIs(
|
|
cached_audio, tts.initial_greeting_audio("Olá, como posso ajudar?")
|
|
)
|
|
|
|
runtime_logger.return_value.info.assert_called_once_with(
|
|
"INITIAL_GREETING_AUDIO_CACHE_HIT | key=%s", "cache-key"
|
|
)
|
|
cache.discard.assert_not_called()
|
|
|
|
async def test_connect_exposes_a_deterministic_prewarm_operation(self) -> None:
|
|
tts = OraclexAITTS(api_key="key-123")
|
|
with mock.patch.object(tts, "_current_connection", new=mock.AsyncMock()) as current_connection:
|
|
await tts.connect(1.5)
|
|
|
|
current_connection.assert_awaited_once_with(1.5)
|
|
await tts.aclose()
|
|
|
|
async def test_stream_merges_text_chunks_into_one_provider_turn(self) -> None:
|
|
greeting_chunks = (
|
|
"Olá! Eu sou a Especialista em Contas e vou ajudar você a entender a sua fatura. ",
|
|
"Posso explicar valores, detalhar serviços e itens eventuais, identificar cobranças que você não reconhece e, se for o caso, realizar ajustes necessários ou solicitações relacionadas à sua conta. ",
|
|
"Então vamos lá, me conte o que você gostaria de entender ou resolver na sua conta.",
|
|
)
|
|
greeting = "".join(greeting_chunks)
|
|
pcm = b"\x01\x00" * 4800
|
|
websocket = _FakeWebSocket(
|
|
[
|
|
_event("audio.clear"),
|
|
_event("audio.delta", delta=base64.b64encode(pcm).decode()),
|
|
_event("audio.done", trace_id="trace-chunked"),
|
|
]
|
|
)
|
|
session = SimpleNamespace(
|
|
ws_connect=mock.AsyncMock(return_value=websocket),
|
|
)
|
|
cache_key = SimpleNamespace(digest="greeting-key", text=greeting)
|
|
cache = mock.Mock()
|
|
cache.key_for.return_value = cache_key
|
|
cache.has.return_value = False
|
|
cache.store_pcm = mock.AsyncMock()
|
|
|
|
provider = OraclexAITTS(
|
|
api_key="key-123",
|
|
websocket_url="wss://example.test/tts",
|
|
http_session=session,
|
|
initial_greeting_audio_cache=cache,
|
|
initial_greeting_agent="contas",
|
|
)
|
|
|
|
self.assertIsNone(provider.initial_greeting_audio(greeting))
|
|
|
|
stream = provider.stream()
|
|
for chunk in greeting_chunks:
|
|
stream.push_text(chunk)
|
|
stream.end_input()
|
|
try:
|
|
audio_events = [event async for event in stream]
|
|
finally:
|
|
await stream.aclose()
|
|
await provider.aclose()
|
|
|
|
self.assertIsInstance(stream, xai_tts_module.SynthesizeStream)
|
|
self.assertEqual(b"".join(event.frame.data.tobytes() for event in audio_events), pcm)
|
|
self.assertEqual(
|
|
[message["type"] for message in websocket.sent],
|
|
["text.clear", "text.delta", "text.done"],
|
|
)
|
|
self.assertEqual(websocket.sent[1]["delta"], greeting)
|
|
cache.store_pcm.assert_awaited_once_with(cache_key, pcm)
|
|
|
|
def test_text_sanitization_is_preserved(self) -> None:
|
|
self.assertEqual(
|
|
xai_tts_module._sanitize_tts_text(
|
|
"TIM_GAMES_KIDS_MES custa R$ 14,99 no dia 29/05/26; a/b?"
|
|
),
|
|
"TIM GAMES KIDS MES custa R 14,99 no dia 29/05/26; a ou b?",
|
|
)
|
|
|
|
async def test_turn_requires_clear_ack_before_emitting_audio(self) -> None:
|
|
payload = base64.b64encode(b"novo").decode()
|
|
connection, emitter, result = await self._synthesize(
|
|
[_event("audio.clear"), _event("audio.delta", delta=payload), _event("audio.done", trace_id="trace-1")]
|
|
)
|
|
|
|
self.assertEqual(bytes(emitter.audio), b"novo")
|
|
self.assertEqual(
|
|
[message["type"] for message in connection._ws.sent],
|
|
["text.clear", "text.delta", "text.done"],
|
|
)
|
|
self.assertEqual(result.trace_id, "trace-1")
|
|
self.assertEqual(result.timing.discarded_messages, [])
|
|
|
|
async def test_residual_audio_is_discarded_before_clear_ack(self) -> None:
|
|
old_payload = base64.b64encode(b"velho").decode()
|
|
new_payload = base64.b64encode(b"novo").decode()
|
|
_connection, emitter, result = await self._synthesize(
|
|
[
|
|
_event("audio.delta", delta=old_payload),
|
|
_event("audio.done"),
|
|
_event("audio.clear"),
|
|
_event("audio.delta", delta=new_payload),
|
|
_event("audio.done", trace_id="trace-new"),
|
|
]
|
|
)
|
|
|
|
self.assertEqual(bytes(emitter.audio), b"novo")
|
|
self.assertEqual(result.trace_id, "trace-new")
|
|
self.assertEqual(result.timing.discarded_messages, ["audio.delta", "audio.done"])
|
|
self.assertEqual(result.timing.clear_discarded_message_count, 2)
|
|
self.assertEqual(result.timing.clear_discarded_audio_bytes, len(b"velho"))
|
|
|
|
async def test_unexpected_clear_after_boundary_fails_and_retires_socket(self) -> None:
|
|
payload = base64.b64encode(b"parcial").decode()
|
|
connection = self._connection(
|
|
[_event("audio.clear"), _event("audio.delta", delta=payload), _event("audio.clear")]
|
|
)
|
|
emitter = _Emitter()
|
|
|
|
with self.assertRaisesRegex(xai_tts_module._XAIPartialAudioFailure, "unexpected_audio_clear"):
|
|
await connection.synthesize_turn(
|
|
"fala",
|
|
output_emitter=emitter,
|
|
stream=_Stream(),
|
|
timeout=1.0,
|
|
turn_index=0,
|
|
connection_reused=True,
|
|
)
|
|
|
|
self.assertEqual(bytes(emitter.audio), b"parcial")
|
|
self.assertIsNone(connection._ws)
|
|
|
|
async def test_first_frame_timeout_is_configurable(self) -> None:
|
|
connection = self._connection([_event("audio.clear")])
|
|
connection._session = SimpleNamespace(
|
|
ws_connect=mock.AsyncMock(return_value=_FakeWebSocket([]))
|
|
)
|
|
emitter = _Emitter()
|
|
with mock.patch.dict(os.environ, {"TTS_FIRST_FRAME_TIMEOUT_S": "0.01"}, clear=False):
|
|
with self.assertRaisesRegex(xai_tts_module.APIConnectionError, "timed out before audio"):
|
|
await connection.synthesize_turn(
|
|
"fala",
|
|
output_emitter=emitter,
|
|
stream=_Stream(),
|
|
timeout=1.0,
|
|
turn_index=0,
|
|
connection_reused=True,
|
|
)
|
|
|
|
self.assertIsNone(connection._ws)
|
|
|
|
async def test_audio_done_without_pcm_retries_once_on_same_socket(self) -> None:
|
|
pcm = b"\x01\x00" * 480
|
|
connection, emitter, result = await self._synthesize([
|
|
_event("audio.clear"),
|
|
_event("audio.done"),
|
|
_event("audio.clear"),
|
|
_event("audio.delta", delta=base64.b64encode(pcm).decode()),
|
|
_event("audio.done", trace_id="trace-after-empty"),
|
|
])
|
|
self.assertEqual(bytes(emitter.audio), pcm)
|
|
self.assertEqual(result.timing.attempts, 2)
|
|
self.assertFalse(connection._ws.closed)
|
|
self.assertEqual(
|
|
[item["type"] for item in connection._ws.sent],
|
|
["text.clear", "text.delta", "text.done"] * 2,
|
|
)
|
|
|
|
async def test_partial_audio_resync_discards_socket_without_replay(self) -> None:
|
|
pcm = b"\x01\x00" * 480
|
|
owner = _EventOwner()
|
|
connection = self._connection([
|
|
_event("audio.clear"),
|
|
_event("audio.delta", delta=base64.b64encode(pcm).decode()),
|
|
], owner=owner)
|
|
emitter = _Emitter()
|
|
websocket = connection._ws
|
|
with self.assertRaisesRegex(xai_tts_module._XAIPartialAudioFailure, "socket_resynchronized=0"):
|
|
await connection.synthesize_turn("fala", output_emitter=emitter, stream=_Stream(), timeout=1.0, turn_index=0, connection_reused=True)
|
|
self.assertEqual([item["type"] for item in websocket.sent].count("text.delta"), 1)
|
|
self.assertEqual([item["type"] for item in websocket.sent].count("text.clear"), 2)
|
|
self.assertEqual(len(owner.events), 1)
|
|
event_name, event = owner.events[0]
|
|
self.assertEqual(event_name, "xai_tts_turn_failed")
|
|
self.assertEqual(event["segment_id"], "segment-test")
|
|
self.assertEqual(event["reason"], "underflow_error")
|
|
self.assertEqual(event["xai_micro_underflows"], 1)
|
|
self.assertGreaterEqual(event["max_playout_underrun_0ms"], 10)
|
|
self.assertEqual(event["attempts"], 1)
|
|
self.assertGreater(event["pcm_bytes"], 0)
|
|
self.assertGreater(event["pcm_duration_ms"], 0)
|
|
self.assertTrue(event["connection_reused"])
|
|
self.assertFalse(event["reconnected"])
|
|
self.assertIsNone(connection._ws)
|
|
|
|
async def test_continuous_underflow_discards_socket_without_replaying_text(self) -> None:
|
|
connection = self._connection([
|
|
_event("audio.clear"),
|
|
_event("audio.delta", delta=base64.b64encode(b"\x01\x00").decode()),
|
|
])
|
|
emitter = _Emitter()
|
|
websocket = connection._ws
|
|
with mock.patch.dict(os.environ, {"TTS_UNDERFLOW_ERROR_MS": "10", "TTS_FIRST_FRAME_TIMEOUT_S": "0.01"}, clear=False):
|
|
with self.assertRaisesRegex(xai_tts_module._XAIPartialAudioFailure, "underflow_error") as raised:
|
|
await connection.synthesize_turn("fala", output_emitter=emitter, stream=_Stream(), timeout=1.0, turn_index=0, connection_reused=True)
|
|
self.assertIn("xai_micro_underflows=1", str(raised.exception))
|
|
self.assertIn("xai_avg_underrun_ms=", str(raised.exception))
|
|
self.assertEqual([item["type"] for item in websocket.sent].count("text.delta"), 1)
|
|
self.assertIsNone(connection._ws)
|
|
|
|
|
|
if __name__ == "__main__":
|
|
unittest.main()
|