Refactor LMNTTTSService to make a websocket connection directly, then use the WebsocketService base class

This commit is contained in:
Mark Backman
2025-01-10 12:34:14 -05:00
parent 5e5de618f3
commit e60a59434f

View File

@@ -4,11 +4,10 @@
# SPDX-License-Identifier: BSD 2-Clause License # SPDX-License-Identifier: BSD 2-Clause License
# #
import asyncio import json
from typing import AsyncGenerator from typing import AsyncGenerator
from loguru import logger from loguru import logger
from tenacity import AsyncRetrying, RetryCallState, stop_after_attempt, wait_exponential
from pipecat.frames.frames import ( from pipecat.frames.frames import (
CancelFrame, CancelFrame,
@@ -23,11 +22,12 @@ from pipecat.frames.frames import (
) )
from pipecat.processors.frame_processor import FrameDirection from pipecat.processors.frame_processor import FrameDirection
from pipecat.services.ai_services import TTSService from pipecat.services.ai_services import TTSService
from pipecat.services.websocket_service import WebsocketService
from pipecat.transcriptions.language import Language from pipecat.transcriptions.language import Language
# See .env.example for LMNT configuration needed # See .env.example for LMNT configuration needed
try: try:
from lmnt.api import Speech import websockets
except ModuleNotFoundError as e: except ModuleNotFoundError as e:
logger.error(f"Exception: {e}") logger.error(f"Exception: {e}")
logger.error( logger.error(
@@ -60,7 +60,7 @@ def language_to_lmnt_language(language: Language) -> str | None:
return result return result
class LmntTTSService(TTSService): class LmntTTSService(TTSService, WebsocketService):
def __init__( def __init__(
self, self,
*, *,
@@ -70,27 +70,21 @@ class LmntTTSService(TTSService):
language: Language = Language.EN, language: Language = Language.EN,
**kwargs, **kwargs,
): ):
# Let TTSService produce TTSStoppedFrames after a short delay of TTSService.__init__(
# no activity. self,
super().__init__(push_stop_frames=True, sample_rate=sample_rate, **kwargs) push_stop_frames=True,
sample_rate=sample_rate,
**kwargs,
)
WebsocketService.__init__(self)
self._api_key = api_key self._api_key = api_key
self._voice_id = voice_id
self._settings = { self._settings = {
"output_format": { "sample_rate": sample_rate,
"container": "raw",
"encoding": "pcm_s16le",
"sample_rate": sample_rate,
},
"language": self.language_to_service_language(language), "language": self.language_to_service_language(language),
"format": "raw", # Use raw format for direct PCM data
} }
self.set_voice(voice_id)
self._speech = None
self._connection = None
self._receive_task = None
# Indicates if we have sent TTSStartedFrame. It will reset to False when
# there's an interruption or TTSStoppedFrame.
self._started = False self._started = False
def can_generate_metrics(self) -> bool: def can_generate_metrics(self) -> bool:
@@ -117,106 +111,105 @@ class LmntTTSService(TTSService):
self._started = False self._started = False
async def _connect(self): async def _connect(self):
await self._connect_lmnt() await self._connect_websocket()
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())
async def _disconnect(self): async def _disconnect(self):
await self._disconnect_lmnt() await self._disconnect_websocket()
if self._receive_task: if self._receive_task:
self._receive_task.cancel() self._receive_task.cancel()
await self._receive_task await self._receive_task
self._receive_task = None self._receive_task = None
async def _connect_lmnt(self): async def _connect_websocket(self):
"""Connect to LMNT websocket."""
try: try:
logger.debug("Connecting to LMNT") logger.debug("Connecting to LMNT")
self._speech = Speech() # Build initial connection message
self._connection = await self._speech.synthesize_streaming( init_msg = {
self._voice_id, "X-API-Key": self._api_key,
format="raw", "voice": self._voice_id,
sample_rate=self._settings["output_format"]["sample_rate"], "format": self._settings["format"],
language=self._settings["language"], "sample_rate": self._settings["sample_rate"],
) "language": self._settings["language"],
}
# Connect to LMNT's websocket directly
self._websocket = await websockets.connect("wss://api.lmnt.com/v1/ai/speech/stream")
# Send initialization message
await self._websocket.send(json.dumps(init_msg))
except Exception as e: except Exception as e:
logger.error(f"{self} initialization error: {e}") logger.error(f"{self} initialization error: {e}")
self._connection = None self._websocket = None
async def _disconnect_lmnt(self): async def _disconnect_websocket(self):
"""Disconnect from LMNT websocket."""
try: try:
await self.stop_all_metrics() await self.stop_all_metrics()
if self._connection: if self._websocket:
logger.debug("Disconnecting from LMNT") logger.debug("Disconnecting from LMNT")
await self._connection.socket.close() # Send EOF message before closing
self._connection = None await self._websocket.send(json.dumps({"eof": True}))
if self._speech: await self._websocket.close()
await self._speech.close() self._websocket = None
self._speech = None
self._started = False self._started = False
except Exception as e: except Exception as e:
logger.error(f"{self} error closing connection: {e}") logger.error(f"{self} error closing websocket: {e}")
def _get_websocket(self):
if self._websocket:
return self._websocket
raise Exception("Websocket not connected")
async def _receive_messages(self): async def _receive_messages(self):
async for msg in self._connection: """Receive messages from LMNT websocket."""
if "error" in msg: async for message in self._get_websocket():
logger.error(f'{self} error: {msg["error"]}') if isinstance(message, bytes):
await self.push_frame(TTSStoppedFrame()) # Raw audio data
await self.stop_all_metrics()
await self.push_error(ErrorFrame(f'{self} error: {msg["error"]}'))
elif "audio" in msg:
await self.stop_ttfb_metrics() await self.stop_ttfb_metrics()
frame = TTSAudioRawFrame( frame = TTSAudioRawFrame(
audio=msg["audio"], audio=message,
sample_rate=self._settings["output_format"]["sample_rate"], sample_rate=self._settings["sample_rate"],
num_channels=1, num_channels=1,
) )
await self.push_frame(frame) await self.push_frame(frame)
else: else:
logger.error(f"{self}: LMNT error, unknown message type: {msg}") try:
msg = json.loads(message)
async def _reconnect_websocket(self, retry_state: RetryCallState): if "error" in msg:
logger.warning(f"{self} reconnecting (attempt: {retry_state.attempt_number})") logger.error(f'{self} error: {msg["error"]}')
await self._disconnect_lmnt() await self.push_frame(TTSStoppedFrame())
await self._connect_lmnt() await self.stop_all_metrics()
await self.push_error(ErrorFrame(f'{self} error: {msg["error"]}'))
async def _receive_task_handler(self): return
while True: except json.JSONDecodeError:
try: logger.error(f"Invalid JSON message: {message}")
async for attempt in AsyncRetrying(
stop=stop_after_attempt(3),
wait=wait_exponential(multiplier=1, min=4, max=10),
before_sleep=self._reconnect_websocket,
reraise=True,
):
with attempt:
await self._receive_messages()
except asyncio.CancelledError:
break
except Exception as e:
message = f"{self} error receiving messages: {e}"
logger.error(message)
await self.push_error(ErrorFrame(message, fatal=True))
break
async def run_tts(self, text: str) -> AsyncGenerator[Frame, None]: async def run_tts(self, text: str) -> AsyncGenerator[Frame, None]:
"""Generate TTS audio from text."""
logger.debug(f"Generating TTS: [{text}]") logger.debug(f"Generating TTS: [{text}]")
try: try:
if not self._connection: if not self._websocket:
await self._connect() await self._connect()
if not self._started:
await self.start_ttfb_metrics()
yield TTSStartedFrame()
self._started = True
try: try:
await self._connection.append_text(text) if not self._started:
await self._connection.flush() await self.start_ttfb_metrics()
yield TTSStartedFrame()
self._started = True
# Send text to LMNT
await self._get_websocket().send(json.dumps({"text": text}))
# Force synthesis
await self._get_websocket().send(json.dumps({"flush": True}))
await self.start_tts_usage_metrics(text) await self.start_tts_usage_metrics(text)
except Exception as e: except Exception as e:
logger.error(f"{self} error sending message: {e}") logger.error(f"{self} error sending message: {e}")