Files
2026-08-21 08:37:51 -03:00

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()