Add websocket support for PlayHT
This commit is contained in:
@@ -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):
|
||||||
|
|||||||
@@ -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" ]
|
||||||
|
|||||||
@@ -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}]")
|
||||||
|
|||||||
Reference in New Issue
Block a user