Merge pull request #1607 from pipecat-ai/aleix/fix-websocket-disconnects
services: fix TTS websocket services disconnections
This commit is contained in:
@@ -55,6 +55,9 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0
|
|||||||
|
|
||||||
### Fixed
|
### Fixed
|
||||||
|
|
||||||
|
- Fixed an issue that would cause TTS websocket-based services to not cleanup
|
||||||
|
resources properly when disconnecting.
|
||||||
|
|
||||||
- Fixed a `TavusVideoService` issue that was causing audio choppiness.
|
- Fixed a `TavusVideoService` issue that was causing audio choppiness.
|
||||||
|
|
||||||
- Fixed an issue in `SmallWebRTCTransport` where an error was thrown if the
|
- Fixed an issue in `SmallWebRTCTransport` where an error was thrown if the
|
||||||
|
|||||||
@@ -185,7 +185,8 @@ class CartesiaTTSService(AudioContextWordTTSService):
|
|||||||
|
|
||||||
async def _connect(self):
|
async def _connect(self):
|
||||||
await self._connect_websocket()
|
await self._connect_websocket()
|
||||||
if not self._receive_task:
|
|
||||||
|
if self._websocket and not self._receive_task:
|
||||||
self._receive_task = self.create_task(self._receive_task_handler(self._report_error))
|
self._receive_task = self.create_task(self._receive_task_handler(self._report_error))
|
||||||
|
|
||||||
async def _disconnect(self):
|
async def _disconnect(self):
|
||||||
@@ -197,7 +198,7 @@ class CartesiaTTSService(AudioContextWordTTSService):
|
|||||||
|
|
||||||
async def _connect_websocket(self):
|
async def _connect_websocket(self):
|
||||||
try:
|
try:
|
||||||
if self._websocket:
|
if self._websocket and self._websocket.open:
|
||||||
return
|
return
|
||||||
logger.debug("Connecting to Cartesia")
|
logger.debug("Connecting to Cartesia")
|
||||||
self._websocket = await websockets.connect(
|
self._websocket = await websockets.connect(
|
||||||
@@ -215,11 +216,11 @@ class CartesiaTTSService(AudioContextWordTTSService):
|
|||||||
if self._websocket:
|
if self._websocket:
|
||||||
logger.debug("Disconnecting from Cartesia")
|
logger.debug("Disconnecting from Cartesia")
|
||||||
await self._websocket.close()
|
await self._websocket.close()
|
||||||
self._websocket = 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}")
|
||||||
|
finally:
|
||||||
|
self._context_id = None
|
||||||
|
self._websocket = None
|
||||||
|
|
||||||
def _get_websocket(self):
|
def _get_websocket(self):
|
||||||
if self._websocket:
|
if self._websocket:
|
||||||
@@ -279,7 +280,7 @@ class CartesiaTTSService(AudioContextWordTTSService):
|
|||||||
logger.debug(f"{self}: Generating TTS [{text}]")
|
logger.debug(f"{self}: Generating TTS [{text}]")
|
||||||
|
|
||||||
try:
|
try:
|
||||||
if not self._websocket:
|
if not self._websocket or self._websocket.closed:
|
||||||
await self._connect()
|
await self._connect()
|
||||||
|
|
||||||
if not self._context_id:
|
if not self._context_id:
|
||||||
|
|||||||
@@ -309,10 +309,10 @@ class ElevenLabsTTSService(InterruptibleWordTTSService):
|
|||||||
async def _connect(self):
|
async def _connect(self):
|
||||||
await self._connect_websocket()
|
await self._connect_websocket()
|
||||||
|
|
||||||
if not self._receive_task:
|
if self._websocket and not self._receive_task:
|
||||||
self._receive_task = self.create_task(self._receive_task_handler(self._report_error))
|
self._receive_task = self.create_task(self._receive_task_handler(self._report_error))
|
||||||
|
|
||||||
if not self._keepalive_task:
|
if self._websocket and not self._keepalive_task:
|
||||||
self._keepalive_task = self.create_task(self._keepalive_task_handler())
|
self._keepalive_task = self.create_task(self._keepalive_task_handler())
|
||||||
|
|
||||||
async def _disconnect(self):
|
async def _disconnect(self):
|
||||||
@@ -328,7 +328,7 @@ class ElevenLabsTTSService(InterruptibleWordTTSService):
|
|||||||
|
|
||||||
async def _connect_websocket(self):
|
async def _connect_websocket(self):
|
||||||
try:
|
try:
|
||||||
if self._websocket:
|
if self._websocket and self._websocket.open:
|
||||||
return
|
return
|
||||||
|
|
||||||
logger.debug("Connecting to ElevenLabs")
|
logger.debug("Connecting to ElevenLabs")
|
||||||
@@ -375,11 +375,11 @@ class ElevenLabsTTSService(InterruptibleWordTTSService):
|
|||||||
logger.debug("Disconnecting from ElevenLabs")
|
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._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}")
|
||||||
|
finally:
|
||||||
|
self._started = False
|
||||||
|
self._websocket = None
|
||||||
|
|
||||||
def _get_websocket(self):
|
def _get_websocket(self):
|
||||||
if self._websocket:
|
if self._websocket:
|
||||||
@@ -419,7 +419,7 @@ class ElevenLabsTTSService(InterruptibleWordTTSService):
|
|||||||
logger.debug(f"{self}: Generating TTS [{text}]")
|
logger.debug(f"{self}: Generating TTS [{text}]")
|
||||||
|
|
||||||
try:
|
try:
|
||||||
if not self._websocket:
|
if not self._websocket or self._websocket.closed:
|
||||||
await self._connect()
|
await self._connect()
|
||||||
|
|
||||||
try:
|
try:
|
||||||
|
|||||||
@@ -104,7 +104,8 @@ class FishAudioTTSService(InterruptibleTTSService):
|
|||||||
|
|
||||||
async def _connect(self):
|
async def _connect(self):
|
||||||
await self._connect_websocket()
|
await self._connect_websocket()
|
||||||
if not self._receive_task:
|
|
||||||
|
if self._websocket and not self._receive_task:
|
||||||
self._receive_task = self.create_task(self._receive_task_handler(self._report_error))
|
self._receive_task = self.create_task(self._receive_task_handler(self._report_error))
|
||||||
|
|
||||||
async def _disconnect(self):
|
async def _disconnect(self):
|
||||||
@@ -116,7 +117,7 @@ class FishAudioTTSService(InterruptibleTTSService):
|
|||||||
|
|
||||||
async def _connect_websocket(self):
|
async def _connect_websocket(self):
|
||||||
try:
|
try:
|
||||||
if self._websocket:
|
if self._websocket and self._websocket.open:
|
||||||
return
|
return
|
||||||
|
|
||||||
logger.debug("Connecting to Fish Audio")
|
logger.debug("Connecting to Fish Audio")
|
||||||
@@ -141,16 +142,17 @@ class FishAudioTTSService(InterruptibleTTSService):
|
|||||||
stop_message = {"event": "stop"}
|
stop_message = {"event": "stop"}
|
||||||
await self._websocket.send(ormsgpack.packb(stop_message))
|
await self._websocket.send(ormsgpack.packb(stop_message))
|
||||||
await self._websocket.close()
|
await self._websocket.close()
|
||||||
self._websocket = None
|
|
||||||
self._request_id = None
|
|
||||||
self._started = False
|
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.error(f"Error closing websocket: {e}")
|
logger.error(f"Error closing websocket: {e}")
|
||||||
|
finally:
|
||||||
|
self._request_id = None
|
||||||
|
self._started = False
|
||||||
|
self._websocket = None
|
||||||
|
|
||||||
async def flush_audio(self):
|
async def flush_audio(self):
|
||||||
"""Flush any buffered audio by sending a flush event to Fish Audio."""
|
"""Flush any buffered audio by sending a flush event to Fish Audio."""
|
||||||
logger.trace(f"{self}: Flushing audio buffers")
|
logger.trace(f"{self}: Flushing audio buffers")
|
||||||
if not self._websocket:
|
if not self._websocket or self._websocket.closed:
|
||||||
return
|
return
|
||||||
flush_message = {"event": "flush"}
|
flush_message = {"event": "flush"}
|
||||||
await self._get_websocket().send(ormsgpack.packb(flush_message))
|
await self._get_websocket().send(ormsgpack.packb(flush_message))
|
||||||
|
|||||||
@@ -285,7 +285,7 @@ class GladiaSTTService(STTService):
|
|||||||
settings = self._prepare_settings()
|
settings = self._prepare_settings()
|
||||||
response = await self._setup_gladia(settings)
|
response = await self._setup_gladia(settings)
|
||||||
self._websocket = await websockets.connect(response["url"])
|
self._websocket = await websockets.connect(response["url"])
|
||||||
if not self._receive_task:
|
if self._websocket and not self._receive_task:
|
||||||
self._receive_task = self.create_task(self._receive_task_handler())
|
self._receive_task = self.create_task(self._receive_task_handler())
|
||||||
|
|
||||||
async def stop(self, frame: EndFrame):
|
async def stop(self, frame: EndFrame):
|
||||||
|
|||||||
@@ -109,7 +109,7 @@ class LmntTTSService(InterruptibleTTSService):
|
|||||||
async def _connect(self):
|
async def _connect(self):
|
||||||
await self._connect_websocket()
|
await self._connect_websocket()
|
||||||
|
|
||||||
if not self._receive_task:
|
if self._websocket and not self._receive_task:
|
||||||
self._receive_task = self.create_task(self._receive_task_handler(self._report_error))
|
self._receive_task = self.create_task(self._receive_task_handler(self._report_error))
|
||||||
|
|
||||||
async def _disconnect(self):
|
async def _disconnect(self):
|
||||||
@@ -122,7 +122,7 @@ class LmntTTSService(InterruptibleTTSService):
|
|||||||
async def _connect_websocket(self):
|
async def _connect_websocket(self):
|
||||||
"""Connect to LMNT websocket."""
|
"""Connect to LMNT websocket."""
|
||||||
try:
|
try:
|
||||||
if self._websocket:
|
if self._websocket and self._websocket.open:
|
||||||
return
|
return
|
||||||
|
|
||||||
logger.debug("Connecting to LMNT")
|
logger.debug("Connecting to LMNT")
|
||||||
@@ -158,11 +158,11 @@ class LmntTTSService(InterruptibleTTSService):
|
|||||||
# errors on the websocket, so we just skip it for now.
|
# errors on the websocket, so we just skip it for now.
|
||||||
# await self._websocket.send(json.dumps({"eof": True}))
|
# await self._websocket.send(json.dumps({"eof": True}))
|
||||||
await self._websocket.close()
|
await self._websocket.close()
|
||||||
self._websocket = None
|
|
||||||
|
|
||||||
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}")
|
||||||
|
finally:
|
||||||
|
self._started = False
|
||||||
|
self._websocket = None
|
||||||
|
|
||||||
def _get_websocket(self):
|
def _get_websocket(self):
|
||||||
if self._websocket:
|
if self._websocket:
|
||||||
@@ -170,7 +170,7 @@ class LmntTTSService(InterruptibleTTSService):
|
|||||||
raise Exception("Websocket not connected")
|
raise Exception("Websocket not connected")
|
||||||
|
|
||||||
async def flush_audio(self):
|
async def flush_audio(self):
|
||||||
if not self._websocket:
|
if not self._websocket or self._websocket.closed:
|
||||||
return
|
return
|
||||||
await self._get_websocket().send(json.dumps({"flush": True}))
|
await self._get_websocket().send(json.dumps({"flush": True}))
|
||||||
|
|
||||||
@@ -203,7 +203,7 @@ class LmntTTSService(InterruptibleTTSService):
|
|||||||
logger.debug(f"{self}: Generating TTS [{text}]")
|
logger.debug(f"{self}: Generating TTS [{text}]")
|
||||||
|
|
||||||
try:
|
try:
|
||||||
if not self._websocket:
|
if not self._websocket or self._websocket.closed:
|
||||||
await self._connect()
|
await self._connect()
|
||||||
|
|
||||||
try:
|
try:
|
||||||
|
|||||||
@@ -106,6 +106,9 @@ class NeuphonicTTSService(InterruptibleTTSService):
|
|||||||
self._started = False
|
self._started = False
|
||||||
self._cumulative_time = 0
|
self._cumulative_time = 0
|
||||||
|
|
||||||
|
self._receive_task = None
|
||||||
|
self._keepalive_task = None
|
||||||
|
|
||||||
def can_generate_metrics(self) -> bool:
|
def can_generate_metrics(self) -> bool:
|
||||||
return True
|
return True
|
||||||
|
|
||||||
@@ -159,8 +162,11 @@ class NeuphonicTTSService(InterruptibleTTSService):
|
|||||||
async def _connect(self):
|
async def _connect(self):
|
||||||
await self._connect_websocket()
|
await self._connect_websocket()
|
||||||
|
|
||||||
self._receive_task = self.create_task(self._receive_task_handler(self._report_error))
|
if self._websocket and not self._receive_task:
|
||||||
self._keepalive_task = self.create_task(self._keepalive_task_handler())
|
self._receive_task = self.create_task(self._receive_task_handler(self._report_error))
|
||||||
|
|
||||||
|
if self._websocket and not self._keepalive_task:
|
||||||
|
self._keepalive_task = self.create_task(self._keepalive_task_handler())
|
||||||
|
|
||||||
async def _disconnect(self):
|
async def _disconnect(self):
|
||||||
if self._receive_task:
|
if self._receive_task:
|
||||||
@@ -175,6 +181,9 @@ class NeuphonicTTSService(InterruptibleTTSService):
|
|||||||
|
|
||||||
async def _connect_websocket(self):
|
async def _connect_websocket(self):
|
||||||
try:
|
try:
|
||||||
|
if self._websocket and self._websocket.open:
|
||||||
|
return
|
||||||
|
|
||||||
logger.debug("Connecting to Neuphonic")
|
logger.debug("Connecting to Neuphonic")
|
||||||
|
|
||||||
tts_config = {
|
tts_config = {
|
||||||
@@ -190,7 +199,6 @@ class NeuphonicTTSService(InterruptibleTTSService):
|
|||||||
url = f"{self._url}/speak/{self._settings['lang_code']}?{'&'.join(query_params)}"
|
url = f"{self._url}/speak/{self._settings['lang_code']}?{'&'.join(query_params)}"
|
||||||
|
|
||||||
self._websocket = await websockets.connect(url)
|
self._websocket = await websockets.connect(url)
|
||||||
|
|
||||||
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
|
||||||
@@ -203,11 +211,11 @@ class NeuphonicTTSService(InterruptibleTTSService):
|
|||||||
if self._websocket:
|
if self._websocket:
|
||||||
logger.debug("Disconnecting from Neuphonic")
|
logger.debug("Disconnecting from Neuphonic")
|
||||||
await self._websocket.close()
|
await self._websocket.close()
|
||||||
self._websocket = None
|
|
||||||
|
|
||||||
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}")
|
||||||
|
finally:
|
||||||
|
self._started = False
|
||||||
|
self._websocket = None
|
||||||
|
|
||||||
async def _receive_messages(self):
|
async def _receive_messages(self):
|
||||||
async for message in self._websocket:
|
async for message in self._websocket:
|
||||||
@@ -235,7 +243,7 @@ class NeuphonicTTSService(InterruptibleTTSService):
|
|||||||
logger.debug(f"Generating TTS: [{text}]")
|
logger.debug(f"Generating TTS: [{text}]")
|
||||||
|
|
||||||
try:
|
try:
|
||||||
if not self._websocket:
|
if not self._websocket or self._websocket.closed:
|
||||||
await self._connect()
|
await self._connect()
|
||||||
|
|
||||||
try:
|
try:
|
||||||
|
|||||||
@@ -157,7 +157,7 @@ class PlayHTTTSService(InterruptibleTTSService):
|
|||||||
async def _connect(self):
|
async def _connect(self):
|
||||||
await self._connect_websocket()
|
await self._connect_websocket()
|
||||||
|
|
||||||
if not self._receive_task:
|
if self._websocket and not self._receive_task:
|
||||||
self._receive_task = self.create_task(self._receive_task_handler(self._report_error))
|
self._receive_task = self.create_task(self._receive_task_handler(self._report_error))
|
||||||
|
|
||||||
async def _disconnect(self):
|
async def _disconnect(self):
|
||||||
@@ -169,7 +169,7 @@ class PlayHTTTSService(InterruptibleTTSService):
|
|||||||
|
|
||||||
async def _connect_websocket(self):
|
async def _connect_websocket(self):
|
||||||
try:
|
try:
|
||||||
if self._websocket:
|
if self._websocket and self._websocket.open:
|
||||||
return
|
return
|
||||||
|
|
||||||
logger.debug("Connecting to PlayHT")
|
logger.debug("Connecting to PlayHT")
|
||||||
@@ -197,11 +197,11 @@ class PlayHTTTSService(InterruptibleTTSService):
|
|||||||
if self._websocket:
|
if self._websocket:
|
||||||
logger.debug("Disconnecting from PlayHT")
|
logger.debug("Disconnecting from PlayHT")
|
||||||
await self._websocket.close()
|
await self._websocket.close()
|
||||||
self._websocket = 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}")
|
||||||
|
finally:
|
||||||
|
self._request_id = None
|
||||||
|
self._websocket = None
|
||||||
|
|
||||||
async def _get_websocket_url(self):
|
async def _get_websocket_url(self):
|
||||||
async with aiohttp.ClientSession() as session:
|
async with aiohttp.ClientSession() as session:
|
||||||
|
|||||||
@@ -168,7 +168,7 @@ class RimeTTSService(AudioContextWordTTSService):
|
|||||||
"""Establish websocket connection and start receive task."""
|
"""Establish websocket connection and start receive task."""
|
||||||
await self._connect_websocket()
|
await self._connect_websocket()
|
||||||
|
|
||||||
if not self._receive_task:
|
if self._websocket and not self._receive_task:
|
||||||
self._receive_task = self.create_task(self._receive_task_handler(self._report_error))
|
self._receive_task = self.create_task(self._receive_task_handler(self._report_error))
|
||||||
|
|
||||||
async def _disconnect(self):
|
async def _disconnect(self):
|
||||||
@@ -182,7 +182,7 @@ class RimeTTSService(AudioContextWordTTSService):
|
|||||||
async def _connect_websocket(self):
|
async def _connect_websocket(self):
|
||||||
"""Connect to Rime websocket API with configured settings."""
|
"""Connect to Rime websocket API with configured settings."""
|
||||||
try:
|
try:
|
||||||
if self._websocket:
|
if self._websocket and self._websocket.open:
|
||||||
return
|
return
|
||||||
|
|
||||||
params = "&".join(f"{k}={v}" for k, v in self._settings.items())
|
params = "&".join(f"{k}={v}" for k, v in self._settings.items())
|
||||||
@@ -201,10 +201,11 @@ class RimeTTSService(AudioContextWordTTSService):
|
|||||||
if self._websocket:
|
if self._websocket:
|
||||||
await self._websocket.send(json.dumps(self._build_eos_msg()))
|
await self._websocket.send(json.dumps(self._build_eos_msg()))
|
||||||
await self._websocket.close()
|
await self._websocket.close()
|
||||||
self._websocket = 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}")
|
||||||
|
finally:
|
||||||
|
self._context_id = None
|
||||||
|
self._websocket = None
|
||||||
|
|
||||||
def _get_websocket(self):
|
def _get_websocket(self):
|
||||||
"""Get active websocket connection or raise exception."""
|
"""Get active websocket connection or raise exception."""
|
||||||
@@ -316,7 +317,7 @@ class RimeTTSService(AudioContextWordTTSService):
|
|||||||
"""
|
"""
|
||||||
logger.debug(f"{self}: Generating TTS [{text}]")
|
logger.debug(f"{self}: Generating TTS [{text}]")
|
||||||
try:
|
try:
|
||||||
if not self._websocket:
|
if not self._websocket or self._websocket.closed:
|
||||||
await self._connect()
|
await self._connect()
|
||||||
|
|
||||||
try:
|
try:
|
||||||
|
|||||||
@@ -31,7 +31,7 @@ class WebsocketService(ABC):
|
|||||||
bool: True if connection is verified working, False otherwise
|
bool: True if connection is verified working, False otherwise
|
||||||
"""
|
"""
|
||||||
try:
|
try:
|
||||||
if not self._websocket:
|
if not self._websocket or self._websocket.closed:
|
||||||
return False
|
return False
|
||||||
await self._websocket.ping()
|
await self._websocket.ping()
|
||||||
return True
|
return True
|
||||||
|
|||||||
Reference in New Issue
Block a user