services(cartesia,elevenlabs): close websocket before the receiving task

This commit is contained in:
Aleix Conchillo Flaqué
2024-09-16 23:54:21 -07:00
parent d9d6571c73
commit 20c019ae16
2 changed files with 31 additions and 30 deletions

View File

@@ -136,24 +136,25 @@ class CartesiaTTSService(AsyncWordTTSService):
) )
self._receive_task = self.get_event_loop().create_task(self._receive_task_handler()) 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._websocket = None self._websocket = None
async def _disconnect(self): async def _disconnect(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._websocket: if self._websocket:
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.exception(f"{self} error closing websocket: {e}") logger.error(f"{self} error closing websocket: {e}")
async def _handle_interruption(self, frame: StartInterruptionFrame, direction: FrameDirection): async def _handle_interruption(self, frame: StartInterruptionFrame, direction: FrameDirection):
await super()._handle_interruption(frame, direction) await super()._handle_interruption(frame, direction)
@@ -166,18 +167,18 @@ class CartesiaTTSService(AsyncWordTTSService):
return return
logger.debug("Flushing audio") logger.debug("Flushing audio")
msg = { msg = {
"transcript": "", "transcript": "",
"continue": False, "continue": False,
"context_id": self._context_id, "context_id": self._context_id,
"model_id": self._model_id, "model_id": self._model_id,
"voice": { "voice": {
"mode": "id", "mode": "id",
"id": self._voice_id "id": self._voice_id
}, },
"output_format": self._output_format, "output_format": self._output_format,
"language": self._language, "language": self._language,
"add_timestamps": True, "add_timestamps": True,
} }
await self._websocket.send(json.dumps(msg)) await self._websocket.send(json.dumps(msg))
async def _receive_task_handler(self): async def _receive_task_handler(self):
@@ -217,7 +218,7 @@ class CartesiaTTSService(AsyncWordTTSService):
except asyncio.CancelledError: except asyncio.CancelledError:
pass pass
except Exception as e: except Exception as e:
logger.exception(f"{self} exception: {e}") logger.error(f"{self} exception: {e}")
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}]")
@@ -255,4 +256,4 @@ class CartesiaTTSService(AsyncWordTTSService):
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

@@ -174,13 +174,18 @@ class ElevenLabsTTSService(AsyncWordTTSService):
} }
await self._websocket.send(json.dumps(msg)) await self._websocket.send(json.dumps(msg))
except Exception as e: except Exception as e:
logger.exception(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(self):
try: try:
await self.stop_all_metrics() await self.stop_all_metrics()
if self._websocket:
await self._websocket.send(json.dumps({"text": ""}))
await self._websocket.close()
self._websocket = None
if self._receive_task: if self._receive_task:
self._receive_task.cancel() self._receive_task.cancel()
await self._receive_task await self._receive_task
@@ -191,13 +196,9 @@ class ElevenLabsTTSService(AsyncWordTTSService):
await self._keepalive_task await self._keepalive_task
self._keepalive_task = None self._keepalive_task = None
if self._websocket:
await self._websocket.close()
self._websocket = 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 websocket: {e}")
async def _receive_task_handler(self): async def _receive_task_handler(self):
try: try:
@@ -215,11 +216,10 @@ class ElevenLabsTTSService(AsyncWordTTSService):
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:
pass pass
except Exception as e: except Exception as e:
logger.exception(f"{self} exception: {e}") logger.error(f"{self} exception: {e}")
async def _keepalive_task_handler(self): async def _keepalive_task_handler(self):
while True: while True:
@@ -229,7 +229,7 @@ class ElevenLabsTTSService(AsyncWordTTSService):
except asyncio.CancelledError: except asyncio.CancelledError:
break break
except Exception as e: except Exception as e:
logger.exception(f"{self} exception: {e}") logger.error(f"{self} exception: {e}")
async def _send_text(self, text: str): async def _send_text(self, text: str):
if self._websocket: if self._websocket:
@@ -260,4 +260,4 @@ class ElevenLabsTTSService(AsyncWordTTSService):
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}")