Add websocket support for PlayHT

This commit is contained in:
Mark Backman
2024-10-17 12:42:26 -04:00
parent 45606e177c
commit da3810f1a2
3 changed files with 279 additions and 62 deletions

View File

@@ -23,9 +23,9 @@ from pipecat.processors.aggregators.llm_response import (
) )
from pipecat.services.openai import OpenAILLMService from pipecat.services.openai import OpenAILLMService
from pipecat.services.playht import PlayHTTTSService from pipecat.services.playht import PlayHTTTSService
from pipecat.transcriptions.language import Language
from pipecat.transports.services.daily import DailyParams, DailyTransport from pipecat.transports.services.daily import DailyParams, DailyTransport
from pipecat.vad.silero import SileroVADAnalyzer from pipecat.vad.silero import SileroVADAnalyzer
from pipecat.transcriptions.language import Language
load_dotenv(override=True) load_dotenv(override=True)
@@ -80,7 +80,15 @@ async def main():
] ]
) )
task = PipelineTask(pipeline, PipelineParams(allow_interruptions=True)) task = PipelineTask(
pipeline,
PipelineParams(
allow_interruptions=True,
enable_metrics=True,
enable_usage_metrics=True,
report_only_initial_ttfb=True,
),
)
@transport.event_handler("on_first_participant_joined") @transport.event_handler("on_first_participant_joined")
async def on_first_participant_joined(transport, participant): async def on_first_participant_joined(transport, participant):

View File

@@ -56,7 +56,7 @@ local = [ "pyaudio~=0.2.14" ]
moondream = [ "einops~=0.8.0", "timm~=1.0.8", "transformers~=4.44.0" ] moondream = [ "einops~=0.8.0", "timm~=1.0.8", "transformers~=4.44.0" ]
openai = [ "openai~=1.50.2", "websockets~=13.1", "python-deepcompare~=1.0.1" ] openai = [ "openai~=1.50.2", "websockets~=13.1", "python-deepcompare~=1.0.1" ]
openpipe = [ "openpipe~=4.24.0" ] openpipe = [ "openpipe~=4.24.0" ]
playht = [ "pyht~=0.1.4" ] playht = [ "pyht~=0.1.4", "websockets~=13.1" ]
silero = [ "onnxruntime>=1.16.1" ] silero = [ "onnxruntime>=1.16.1" ]
together = [ "openai~=1.50.2" ] together = [ "openai~=1.50.2" ]
websocket = [ "websockets~=13.1", "fastapi~=0.115.0" ] websocket = [ "websockets~=13.1", "fastapi~=0.115.0" ]

View File

@@ -4,19 +4,29 @@
# SPDX-License-Identifier: BSD 2-Clause License # SPDX-License-Identifier: BSD 2-Clause License
# #
import asyncio
import io import io
import json
import struct import struct
from typing import AsyncGenerator, Optional from typing import AsyncGenerator, Optional
import aiohttp
import websockets
from loguru import logger from loguru import logger
from pydantic.main import BaseModel from pydantic.main import BaseModel
from pipecat.frames.frames import ( from pipecat.frames.frames import (
CancelFrame,
EndFrame,
ErrorFrame,
Frame, Frame,
StartFrame,
StartInterruptionFrame,
TTSAudioRawFrame, TTSAudioRawFrame,
TTSStartedFrame, TTSStartedFrame,
TTSStoppedFrame, TTSStoppedFrame,
) )
from pipecat.processors.frame_processor import FrameDirection
from pipecat.services.ai_services import TTSService from pipecat.services.ai_services import TTSService
from pipecat.transcriptions.language import Language from pipecat.transcriptions.language import Language
@@ -32,12 +42,267 @@ except ModuleNotFoundError as e:
raise Exception(f"Missing module: {e}") raise Exception(f"Missing module: {e}")
def language_to_playht_language(language: Language) -> str | None:
match language:
case Language.BG:
return "BULGARIAN"
case Language.CA:
return "CATALAN"
case Language.CS:
return "CZECH"
case Language.DA:
return "DANISH"
case Language.DE:
return "GERMAN"
case (
Language.EN
| Language.EN_US
| Language.EN_GB
| Language.EN_AU
| Language.EN_NZ
| Language.EN_IN
):
return "ENGLISH"
case Language.ES:
return "SPANISH"
case Language.FR | Language.FR_CA:
return "FRENCH"
case Language.EL:
return "GREEK"
case Language.HI:
return "HINDI"
case Language.HU:
return "HUNGARIAN"
case Language.ID:
return "INDONESIAN"
case Language.IT:
return "ITALIAN"
case Language.JA:
return "JAPANESE"
case Language.KO:
return "KOREAN"
case Language.MS:
return "MALAY"
case Language.NL:
return "DUTCH"
case Language.PL:
return "POLISH"
case Language.PT | Language.PT_BR:
return "PORTUGUESE"
case Language.RU:
return "RUSSIAN"
case Language.SV:
return "SWEDISH"
case Language.TH:
return "THAI"
case Language.TR:
return "TURKISH"
case Language.UK:
return "UKRAINIAN"
return None
class PlayHTTTSService(TTSService): class PlayHTTTSService(TTSService):
class InputParams(BaseModel): class InputParams(BaseModel):
language: Optional[Language] = Language.EN language: Optional[Language] = Language.EN
speed: Optional[float] = 1.0 speed: Optional[float] = 1.0
seed: Optional[int] = None seed: Optional[int] = None
def __init__(
self,
*,
api_key: str,
user_id: str,
voice_url: str,
voice_engine: str = "PlayHT3.0-mini",
sample_rate: int = 16000,
output_format: str = "wav",
params: InputParams = InputParams(),
**kwargs,
):
super().__init__(sample_rate=sample_rate, **kwargs)
self._api_key = api_key
self._user_id = user_id
self._websocket_url = None
self._websocket = None
self._receive_task = None
self._settings = {
"sample_rate": sample_rate,
"language": self.language_to_service_language(params.language)
if params.language
else Language.EN,
"output_format": output_format,
"voice_engine": voice_engine,
"speed": params.speed,
"seed": params.seed,
}
self.set_model_name(voice_engine)
self.set_voice(voice_url)
def can_generate_metrics(self) -> bool:
return True
def language_to_service_language(self, language: Language) -> str | None:
# Keep your existing language mapping logic here
pass
async def start(self, frame: StartFrame):
await super().start(frame)
await self._connect()
async def stop(self, frame: EndFrame):
await super().stop(frame)
await self._disconnect()
async def cancel(self, frame: CancelFrame):
await super().cancel(frame)
await self._disconnect()
async def _connect(self):
try:
if not self._websocket_url:
await self._get_websocket_url()
if not isinstance(self._websocket_url, str):
raise ValueError("WebSocket URL is not a string")
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:
logger.error(f"{self} initialization error: {ve}")
self._websocket = None
except Exception as e:
logger.error(f"{self} initialization error: {e}")
self._websocket = None
async def _disconnect(self):
try:
await self.stop_all_metrics()
if self._websocket:
await self._websocket.close()
self._websocket = None
if self._receive_task:
self._receive_task.cancel()
await self._receive_task
self._receive_task = None
except Exception as e:
logger.error(f"{self} error closing websocket: {e}")
async def _get_websocket_url(self):
async with aiohttp.ClientSession() as session:
async with session.post(
"https://api.play.ht/api/v3/websocket-auth",
headers={
"Authorization": f"Bearer {self._api_key}",
"X-User-Id": self._user_id,
"Content-Type": "application/json",
},
) as response:
if response.status in (200, 201):
data = await response.json()
if "websocket_url" in data and isinstance(data["websocket_url"], str):
self._websocket_url = data["websocket_url"]
else:
raise ValueError("Invalid or missing WebSocket URL in response")
else:
raise Exception(f"Failed to get WebSocket URL: {response.status}")
def _get_websocket(self):
if self._websocket:
return self._websocket
raise Exception("Websocket not connected")
async def _handle_interruption(self, frame: StartInterruptionFrame, direction: FrameDirection):
await super()._handle_interruption(frame, direction)
await self.stop_all_metrics()
async def _receive_task_handler(self):
try:
header_size = 78 # Size of the WAV header + extra bytes we want to skip
header_received = False
async for message in self._get_websocket():
if isinstance(message, bytes):
chunk_size = len(message)
# Skip the WAV header
if not header_received and chunk_size == header_size:
header_received = True
continue
await self.stop_ttfb_metrics()
frame = TTSAudioRawFrame(message, self._settings["sample_rate"], 1)
await self.push_frame(frame)
else:
logger.debug(f"Received text message: {message}")
try:
msg = json.loads(message)
if "request_id" in msg:
await self.push_frame(TTSStoppedFrame())
header_received = False # Reset for the next audio stream
elif "error" in msg:
logger.error(f"{self} error: {msg}")
await self.push_error(ErrorFrame(f'{self} error: {msg["error"]}'))
except json.JSONDecodeError:
logger.error(f"Invalid JSON message: {message}")
except asyncio.CancelledError:
pass
except Exception as e:
logger.error(f"{self} exception in receive task: {e}")
async def run_tts(self, text: str) -> AsyncGenerator[Frame, None]:
logger.debug(f"Generating TTS: [{text}]")
try:
# Reconnect if the websocket is closed
if not self._websocket or self._websocket.closed:
await self._connect()
await self.start_ttfb_metrics()
yield TTSStartedFrame()
tts_command = {
"text": text,
"voice": self._voice_id,
"voice_engine": self._settings["voice_engine"],
"output_format": self._settings["output_format"],
"sample_rate": self._settings["sample_rate"],
"language": self._settings["language"],
"speed": self._settings["speed"],
"seed": self._settings["seed"],
}
try:
await self._get_websocket().send(json.dumps(tts_command))
await self.start_tts_usage_metrics(text)
except Exception as e:
logger.error(f"{self} error sending message: {e}")
yield TTSStoppedFrame()
await self._disconnect()
await self._connect()
return
# The actual audio frames will be handled in _receive_task_handler
yield None
except Exception as e:
logger.error(f"{self} error generating TTS: {e}")
yield ErrorFrame(f"{self} error: {str(e)}")
finally:
await self.stop_all_metrics()
class PlayHTHttpTTSService(TTSService):
class InputParams(BaseModel):
language: Optional[Language] = Language.EN
speed: Optional[float] = 1.0
seed: Optional[int] = None
def __init__( def __init__(
self, self,
*, *,
@@ -52,11 +317,11 @@ class PlayHTTTSService(TTSService):
super().__init__(sample_rate=sample_rate, **kwargs) super().__init__(sample_rate=sample_rate, **kwargs)
self._user_id = user_id self._user_id = user_id
self._speech_key = api_key self._api_key = api_key
self._client = AsyncClient( self._client = AsyncClient(
user_id=self._user_id, user_id=self._user_id,
api_key=self._speech_key, api_key=self._api_key,
) )
self._settings = { self._settings = {
"sample_rate": sample_rate, "sample_rate": sample_rate,
@@ -83,63 +348,7 @@ class PlayHTTTSService(TTSService):
return True return True
def language_to_service_language(self, language: Language) -> str | None: def language_to_service_language(self, language: Language) -> str | None:
match language: return language_to_playht_language(language)
case Language.BG:
return "BULGARIAN"
case Language.CA:
return "CATALAN"
case Language.CS:
return "CZECH"
case Language.DA:
return "DANISH"
case Language.DE:
return "GERMAN"
case (
Language.EN
| Language.EN_US
| Language.EN_GB
| Language.EN_AU
| Language.EN_NZ
| Language.EN_IN
):
return "ENGLISH"
case Language.ES:
return "SPANISH"
case Language.FR | Language.FR_CA:
return "FRENCH"
case Language.EL:
return "GREEK"
case Language.HI:
return "HINDI"
case Language.HU:
return "HUNGARIAN"
case Language.ID:
return "INDONESIAN"
case Language.IT:
return "ITALIAN"
case Language.JA:
return "JAPANESE"
case Language.KO:
return "KOREAN"
case Language.MS:
return "MALAY"
case Language.NL:
return "DUTCH"
case Language.PL:
return "POLISH"
case Language.PT | Language.PT_BR:
return "PORTUGUESE"
case Language.RU:
return "RUSSIAN"
case Language.SV:
return "SWEDISH"
case Language.TH:
return "THAI"
case Language.TR:
return "TURKISH"
case Language.UK:
return "UKRAINIAN"
return None
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}]")