Refactor LMNTTTSService to make a websocket connection directly, then use the WebsocketService base class
This commit is contained in:
@@ -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}")
|
||||||
|
|||||||
Reference in New Issue
Block a user