services: improve Cartesia, 11Labs, PlayHT and LMNT TTS reconnection

This commit is contained in:
Aleix Conchillo Flaqué
2024-12-06 09:25:04 -08:00
parent b05809be2e
commit bafb867ffc
6 changed files with 204 additions and 147 deletions

View File

@@ -68,6 +68,9 @@ async def on_audio_data(processor, audio, sample_rate, num_channels):
### Fixed ### Fixed
- Fixed Cartesia, ElevenLabs, LMNT and PlayHT TTS websocket
reconnection. Before, if an error occurred no reconnection was happening.
- Fixed a `BaseOutputTransport` issue that was causing audio to be discarded - Fixed a `BaseOutputTransport` issue that was causing audio to be discarded
after an `EndFrame` was received. after an `EndFrame` was received.

View File

@@ -184,28 +184,37 @@ class CartesiaTTSService(WordTTSService):
await self._disconnect() await self._disconnect()
async def _connect(self): async def _connect(self):
await self._connect_websocket()
self._receive_task = self.get_event_loop().create_task(self._receive_task_handler())
async def _disconnect(self):
await self._disconnect_websocket()
if self._receive_task:
self._receive_task.cancel()
await self._receive_task
self._receive_task = None
async def _connect_websocket(self):
try: try:
logger.debug("Connecting to Cartesia")
self._websocket = await websockets.connect( self._websocket = await websockets.connect(
f"{self._url}?api_key={self._api_key}&cartesia_version={self._cartesia_version}" f"{self._url}?api_key={self._api_key}&cartesia_version={self._cartesia_version}"
) )
self._receive_task = self.get_event_loop().create_task(self._receive_task_handler())
except Exception as e: except Exception as e:
logger.error(f"{self} initialization error: {e}") logger.error(f"{self} initialization error: {e}")
self._websocket = None self._websocket = None
async def _disconnect(self): async def _disconnect_websocket(self):
try: try:
await self.stop_all_metrics() await self.stop_all_metrics()
if self._websocket: if self._websocket:
logger.debug("Disconnecting from Cartesia")
await self._websocket.close() await self._websocket.close()
self._websocket = None self._websocket = None
if self._receive_task:
self._receive_task.cancel()
await self._receive_task
self._receive_task = None
self._context_id = None self._context_id = None
except Exception as e: except Exception as e:
logger.error(f"{self} error closing websocket: {e}") logger.error(f"{self} error closing websocket: {e}")
@@ -228,44 +237,51 @@ class CartesiaTTSService(WordTTSService):
await self._websocket.send(msg) await self._websocket.send(msg)
async def _receive_task_handler(self): async def _receive_task_handler(self):
try: while True:
async for message in self._get_websocket(): try:
msg = json.loads(message) async for message in self._get_websocket():
if not msg or msg["context_id"] != self._context_id: msg = json.loads(message)
continue if not msg or msg["context_id"] != self._context_id:
if msg["type"] == "done": continue
await self.stop_ttfb_metrics() if msg["type"] == "done":
# Unset _context_id but not the _context_id_start_timestamp await self.stop_ttfb_metrics()
# because we are likely still playing out audio and need the # Unset _context_id but not the _context_id_start_timestamp
# timestamp to set send context frames. # because we are likely still playing out audio and need the
self._context_id = None # timestamp to set send context frames.
await self.add_word_timestamps( self._context_id = None
[("TTSStoppedFrame", 0), ("LLMFullResponseEndFrame", 0), ("Reset", 0)] await self.add_word_timestamps(
) [("TTSStoppedFrame", 0), ("LLMFullResponseEndFrame", 0), ("Reset", 0)]
elif msg["type"] == "timestamps": )
await self.add_word_timestamps( elif msg["type"] == "timestamps":
list(zip(msg["word_timestamps"]["words"], msg["word_timestamps"]["start"])) await self.add_word_timestamps(
) list(
elif msg["type"] == "chunk": zip(
await self.stop_ttfb_metrics() msg["word_timestamps"]["words"], msg["word_timestamps"]["start"]
self.start_word_timestamps() )
frame = TTSAudioRawFrame( )
audio=base64.b64decode(msg["data"]), )
sample_rate=self._settings["output_format"]["sample_rate"], elif msg["type"] == "chunk":
num_channels=1, await self.stop_ttfb_metrics()
) self.start_word_timestamps()
await self.push_frame(frame) frame = TTSAudioRawFrame(
elif msg["type"] == "error": audio=base64.b64decode(msg["data"]),
logger.error(f"{self} error: {msg}") sample_rate=self._settings["output_format"]["sample_rate"],
await self.push_frame(TTSStoppedFrame()) num_channels=1,
await self.stop_all_metrics() )
await self.push_error(ErrorFrame(f'{self} error: {msg["error"]}')) await self.push_frame(frame)
else: elif msg["type"] == "error":
logger.error(f"Cartesia error, unknown message type: {msg}") logger.error(f"{self} error: {msg}")
except asyncio.CancelledError: await self.push_frame(TTSStoppedFrame())
pass await self.stop_all_metrics()
except Exception as e: await self.push_error(ErrorFrame(f'{self} error: {msg["error"]}'))
logger.error(f"{self} exception: {e}") else:
logger.error(f"{self} error, unknown message type: {msg}")
except asyncio.CancelledError:
break
except Exception as e:
logger.error(f"{self} exception: {e}")
await self._disconnect_websocket()
await self._connect_websocket()
async def process_frame(self, frame: Frame, direction: FrameDirection): async def process_frame(self, frame: Frame, direction: FrameDirection):
await super().process_frame(frame, direction) await super().process_frame(frame, direction)
@@ -386,8 +402,6 @@ class CartesiaHttpTTSService(TTSService):
_experimental_voice_controls=voice_controls, _experimental_voice_controls=voice_controls,
) )
await self.stop_ttfb_metrics()
frame = TTSAudioRawFrame( frame = TTSAudioRawFrame(
audio=output["audio"], audio=output["audio"],
sample_rate=self._settings["output_format"]["sample_rate"], sample_rate=self._settings["output_format"]["sample_rate"],
@@ -398,4 +412,6 @@ class CartesiaHttpTTSService(TTSService):
logger.error(f"{self} exception: {e}") logger.error(f"{self} exception: {e}")
await self.start_tts_usage_metrics(text) await self.start_tts_usage_metrics(text)
await self.stop_ttfb_metrics()
yield TTSStoppedFrame() yield TTSStoppedFrame()

View File

@@ -192,15 +192,14 @@ class DeepgramSTTService(STTService):
yield None yield None
async def _connect(self): async def _connect(self):
if await self._connection.start(self._settings): logger.debug("Connecting to Deepgram")
logger.info(f"{self}: Connected to Deepgram") if not await self._connection.start(self._settings):
else: logger.error(f"{self}: unable to connect to Deepgram")
logger.error(f"{self}: Unable to connect to Deepgram")
async def _disconnect(self): async def _disconnect(self):
if self._connection.is_connected: if self._connection.is_connected:
logger.debug("Disconnecting from Deepgram")
await self._connection.finish() await self._connection.finish()
logger.info(f"{self}: Disconnected from Deepgram")
async def _on_speech_started(self, *args, **kwargs): async def _on_speech_started(self, *args, **kwargs):
await self.start_ttfb_metrics() await self.start_ttfb_metrics()

View File

@@ -281,7 +281,28 @@ class ElevenLabsTTSService(WordTTSService):
await self.resume_processing_frames() await self.resume_processing_frames()
async def _connect(self): async def _connect(self):
await self._connect_websocket()
self._receive_task = self.get_event_loop().create_task(self._receive_task_handler())
self._keepalive_task = self.get_event_loop().create_task(self._keepalive_task_handler())
async def _disconnect(self):
if self._receive_task:
self._receive_task.cancel()
await self._receive_task
self._receive_task = None
if self._keepalive_task:
self._keepalive_task.cancel()
await self._keepalive_task
self._keepalive_task = None
await self._disconnect_websocket()
async def _connect_websocket(self):
try: try:
logger.debug("Connecting to ElevenLabs")
voice_id = self._voice_id voice_id = self._voice_id
model = self.model_name model = self.model_name
output_format = self._settings["output_format"] output_format = self._settings["output_format"]
@@ -300,8 +321,6 @@ class ElevenLabsTTSService(WordTTSService):
) )
self._websocket = await websockets.connect(url) self._websocket = await websockets.connect(url)
self._receive_task = self.get_event_loop().create_task(self._receive_task_handler())
self._keepalive_task = self.get_event_loop().create_task(self._keepalive_task_handler())
# According to ElevenLabs, we should always start with a single space. # According to ElevenLabs, we should always start with a single space.
msg: Dict[str, Any] = { msg: Dict[str, Any] = {
@@ -315,49 +334,42 @@ class ElevenLabsTTSService(WordTTSService):
logger.error(f"{self} initialization error: {e}") logger.error(f"{self} initialization error: {e}")
self._websocket = None self._websocket = None
async def _disconnect(self): async def _disconnect_websocket(self):
try: try:
await self.stop_all_metrics() await self.stop_all_metrics()
if self._websocket: if self._websocket:
logger.debug("Disconnecting from ElevenLabs")
await self._websocket.send(json.dumps({"text": ""})) await self._websocket.send(json.dumps({"text": ""}))
await self._websocket.close() await self._websocket.close()
self._websocket = None self._websocket = None
if self._receive_task:
self._receive_task.cancel()
await self._receive_task
self._receive_task = None
if self._keepalive_task:
self._keepalive_task.cancel()
await self._keepalive_task
self._keepalive_task = None
self._started = False self._started = False
except Exception as e: except Exception as e:
logger.error(f"{self} error closing websocket: {e}") logger.error(f"{self} error closing websocket: {e}")
async def _receive_task_handler(self): async def _receive_task_handler(self):
try: while True:
async for message in self._websocket: try:
msg = json.loads(message) async for message in self._websocket:
if msg.get("audio"): msg = json.loads(message)
await self.stop_ttfb_metrics() if msg.get("audio"):
self.start_word_timestamps() await self.stop_ttfb_metrics()
self.start_word_timestamps()
audio = base64.b64decode(msg["audio"]) audio = base64.b64decode(msg["audio"])
frame = TTSAudioRawFrame(audio, self._settings["sample_rate"], 1) frame = TTSAudioRawFrame(audio, self._settings["sample_rate"], 1)
await self.push_frame(frame) await self.push_frame(frame)
if msg.get("alignment"):
if msg.get("alignment"): word_times = calculate_word_times(msg["alignment"], self._cumulative_time)
word_times = calculate_word_times(msg["alignment"], self._cumulative_time) await self.add_word_timestamps(word_times)
await self.add_word_timestamps(word_times) self._cumulative_time = word_times[-1][1]
self._cumulative_time = word_times[-1][1] except asyncio.CancelledError:
except asyncio.CancelledError: break
pass except Exception as e:
except Exception as e: logger.error(f"{self} exception: {e}")
logger.error(f"{self} exception: {e}") await self._disconnect_websocket()
await self._connect_websocket()
async def _keepalive_task_handler(self): async def _keepalive_task_handler(self):
while True: while True:

View File

@@ -116,7 +116,22 @@ class LmntTTSService(TTSService):
self._started = False self._started = False
async def _connect(self): async def _connect(self):
await self._connect_lmnt()
self._receive_task = self.get_event_loop().create_task(self._receive_task_handler())
async def _disconnect(self):
await self._disconnect_lmnt()
if self._receive_task:
self._receive_task.cancel()
await self._receive_task
self._receive_task = None
async def _connect_lmnt(self):
try: try:
logger.debug("Connecting to LMNT")
self._speech = Speech() self._speech = Speech()
self._connection = await self._speech.synthesize_streaming( self._connection = await self._speech.synthesize_streaming(
self._voice_id, self._voice_id,
@@ -124,51 +139,51 @@ class LmntTTSService(TTSService):
sample_rate=self._settings["output_format"]["sample_rate"], sample_rate=self._settings["output_format"]["sample_rate"],
language=self._settings["language"], language=self._settings["language"],
) )
self._receive_task = self.get_event_loop().create_task(self._receive_task_handler())
except Exception as e: except Exception as e:
logger.exception(f"{self} initialization error: {e}") logger.error(f"{self} initialization error: {e}")
self._connection = None self._connection = None
async def _disconnect(self): async def _disconnect_lmnt(self):
try: try:
await self.stop_all_metrics() await self.stop_all_metrics()
if self._receive_task:
self._receive_task.cancel()
await self._receive_task
self._receive_task = None
if self._connection: if self._connection:
logger.debug("Disconnecting from LMNT")
await self._connection.socket.close() await self._connection.socket.close()
self._connection = None self._connection = None
if self._speech: if self._speech:
await self._speech.close() await self._speech.close()
self._speech = None self._speech = None
self._started = False self._started = False
except Exception as e: except Exception as e:
logger.exception(f"{self} error closing websocket: {e}") logger.error(f"{self} error closing connection: {e}")
async def _receive_task_handler(self): async def _receive_task_handler(self):
try: while True:
async for msg in self._connection: try:
if "error" in msg: async for msg in self._connection:
logger.error(f'{self} error: {msg["error"]}') if "error" in msg:
await self.push_frame(TTSStoppedFrame()) logger.error(f'{self} error: {msg["error"]}')
await self.stop_all_metrics() await self.push_frame(TTSStoppedFrame())
await self.push_error(ErrorFrame(f'{self} error: {msg["error"]}')) await self.stop_all_metrics()
elif "audio" in msg: await self.push_error(ErrorFrame(f'{self} error: {msg["error"]}'))
await self.stop_ttfb_metrics() elif "audio" in msg:
frame = TTSAudioRawFrame( await self.stop_ttfb_metrics()
audio=msg["audio"], frame = TTSAudioRawFrame(
sample_rate=self._settings["output_format"]["sample_rate"], audio=msg["audio"],
num_channels=1, sample_rate=self._settings["output_format"]["sample_rate"],
) num_channels=1,
await self.push_frame(frame) )
else: await self.push_frame(frame)
logger.error(f"LMNT error, unknown message type: {msg}") else:
except asyncio.CancelledError: logger.error(f"{self}: LMNT error, unknown message type: {msg}")
pass except asyncio.CancelledError:
except Exception as e: break
logger.exception(f"{self} exception: {e}") except Exception as e:
logger.error(f"{self} exception: {e}")
await self._disconnect_lmnt()
await self._connect_lmnt()
async def run_tts(self, text: str) -> AsyncGenerator[Frame, None]: async def run_tts(self, text: str) -> AsyncGenerator[Frame, None]:
logger.debug(f"Generating TTS: [{text}]") logger.debug(f"Generating TTS: [{text}]")
@@ -194,4 +209,4 @@ class LmntTTSService(TTSService):
return return
yield None yield None
except Exception as e: except Exception as e:
logger.exception(f"{self} exception: {e}") logger.error(f"{self} exception: {e}")

View File

@@ -145,7 +145,22 @@ class PlayHTTTSService(TTSService):
await self._disconnect() await self._disconnect()
async def _connect(self): async def _connect(self):
await self._connect_websocket()
self._receive_task = self.get_event_loop().create_task(self._receive_task_handler())
async def _disconnect(self):
await self._disconnect_websocket()
if self._receive_task:
self._receive_task.cancel()
await self._receive_task
self._receive_task = None
async def _connect_websocket(self):
try: try:
logger.debug("Connecting to PlayHT")
if not self._websocket_url: if not self._websocket_url:
await self._get_websocket_url() await self._get_websocket_url()
@@ -153,8 +168,6 @@ class PlayHTTTSService(TTSService):
raise ValueError("WebSocket URL is not a string") raise ValueError("WebSocket URL is not a string")
self._websocket = await websockets.connect(self._websocket_url) self._websocket = await websockets.connect(self._websocket_url)
self._receive_task = self.get_event_loop().create_task(self._receive_task_handler())
logger.debug("Connected to TTS WebSocket")
except ValueError as ve: except ValueError as ve:
logger.error(f"{self} initialization error: {ve}") logger.error(f"{self} initialization error: {ve}")
self._websocket = None self._websocket = None
@@ -162,19 +175,15 @@ class PlayHTTTSService(TTSService):
logger.error(f"{self} initialization error: {e}") logger.error(f"{self} initialization error: {e}")
self._websocket = None self._websocket = None
async def _disconnect(self): async def _disconnect_websocket(self):
try: try:
await self.stop_all_metrics() await self.stop_all_metrics()
if self._websocket: if self._websocket:
logger.debug("Disconnecting from PlayHT")
await self._websocket.close() await self._websocket.close()
self._websocket = None self._websocket = None
if self._receive_task:
self._receive_task.cancel()
await self._receive_task
self._receive_task = None
self._request_id = None self._request_id = None
except Exception as e: except Exception as e:
logger.error(f"{self} error closing websocket: {e}") logger.error(f"{self} error closing websocket: {e}")
@@ -209,31 +218,34 @@ class PlayHTTTSService(TTSService):
self._request_id = None self._request_id = None
async def _receive_task_handler(self): async def _receive_task_handler(self):
try: while True:
async for message in self._get_websocket(): try:
if isinstance(message, bytes): async for message in self._get_websocket():
# Skip the WAV header message if isinstance(message, bytes):
if message.startswith(b"RIFF"): # Skip the WAV header message
continue if message.startswith(b"RIFF"):
await self.stop_ttfb_metrics() continue
frame = TTSAudioRawFrame(message, self._settings["sample_rate"], 1) await self.stop_ttfb_metrics()
await self.push_frame(frame) frame = TTSAudioRawFrame(message, self._settings["sample_rate"], 1)
else: await self.push_frame(frame)
logger.debug(f"Received text message: {message}") else:
try: logger.debug(f"Received text message: {message}")
msg = json.loads(message) try:
if "request_id" in msg and msg["request_id"] == self._request_id: msg = json.loads(message)
await self.push_frame(TTSStoppedFrame()) if "request_id" in msg and msg["request_id"] == self._request_id:
self._request_id = None await self.push_frame(TTSStoppedFrame())
elif "error" in msg: self._request_id = None
logger.error(f"{self} error: {msg}") elif "error" in msg:
await self.push_error(ErrorFrame(f'{self} error: {msg["error"]}')) logger.error(f"{self} error: {msg}")
except json.JSONDecodeError: await self.push_error(ErrorFrame(f'{self} error: {msg["error"]}'))
logger.error(f"Invalid JSON message: {message}") except json.JSONDecodeError:
except asyncio.CancelledError: logger.error(f"Invalid JSON message: {message}")
pass except asyncio.CancelledError:
except Exception as e: break
logger.error(f"{self} exception in receive task: {e}") except Exception as e:
logger.error(f"{self} exception in receive task: {e}")
await self._disconnect_websocket()
await self._connect_websocket()
async def process_frame(self, frame: Frame, direction: FrameDirection): async def process_frame(self, frame: Frame, direction: FrameDirection):
await super().process_frame(frame, direction) await super().process_frame(frame, direction)
@@ -381,4 +393,4 @@ class PlayHTHttpTTSService(TTSService):
yield frame yield frame
yield TTSStoppedFrame() yield TTSStoppedFrame()
except Exception as e: except Exception as e:
logger.exception(f"{self} error generating TTS: {e}") logger.error(f"{self} error generating TTS: {e}")