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