Merge branch 'main' into recording

This commit is contained in:
Adrian Cowham
2024-10-11 10:36:04 -07:00
23 changed files with 756 additions and 541 deletions

View File

@@ -5,6 +5,32 @@ All notable changes to **Pipecat** will be documented in this file.
The format is based on [Keep a Changelog](https://keepachangelog.com/en/1.0.0/), The format is based on [Keep a Changelog](https://keepachangelog.com/en/1.0.0/),
and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0.html). and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0.html).
## [0.0.43] - 2024-10-10
### Added
- Added a new util called `MarkdownTextFilter` which is a subclass of a new
base class called `BaseTextFilter`. This is a configurable utility which
is intended to filter text received by TTS services.
- Added new `RTVIUserLLMTextProcessor`. This processor will send an RTVI
`user-llm-text` message with the user content's that was sent to the LLM.
### Changed
- `TransportMessageFrame` doesn't have an `urgent` field anymore, instead
there's now a `TransportMessageUrgentFrame` which is a `SystemFrame` and
therefore skip all internal queuing.
- For TTS services, convert inputted languages to match each service's language
format
### Fixed
- Fixed an issue where changing a language with the Deepgram STT service
wouldn't apply the change. This was fixed by disconnecting and reconnecting
when the language changes.
## [0.0.42] - 2024-10-02 ## [0.0.42] - 2024-10-02
### Added ### Added

View File

@@ -21,6 +21,7 @@ classifiers = [
] ]
dependencies = [ dependencies = [
"aiohttp~=3.10.3", "aiohttp~=3.10.3",
"Markdown~=3.7",
"numpy~=1.26.4", "numpy~=1.26.4",
"loguru~=0.7.2", "loguru~=0.7.2",
"Pillow~=10.4.0", "Pillow~=10.4.0",
@@ -40,7 +41,7 @@ azure = [ "azure-cognitiveservices-speech~=1.40.0" ]
canonical = [ "aiofiles~=24.1.0" ] canonical = [ "aiofiles~=24.1.0" ]
cartesia = [ "cartesia~=1.0.13", "websockets~=12.0" ] cartesia = [ "cartesia~=1.0.13", "websockets~=12.0" ]
daily = [ "daily-python~=0.11.0" ] daily = [ "daily-python~=0.11.0" ]
deepgram = [ "deepgram-sdk~=3.5.0" ] deepgram = [ "deepgram-sdk~=3.7.3" ]
elevenlabs = [ "websockets~=12.0" ] elevenlabs = [ "websockets~=12.0" ]
examples = [ "python-dotenv~=1.0.1", "flask~=3.0.3", "flask_cors~=4.0.1" ] examples = [ "python-dotenv~=1.0.1", "flask~=3.0.3", "flask_cors~=4.0.1" ]
fal = [ "fal-client~=0.4.1" ] fal = [ "fal-client~=0.4.1" ]

View File

@@ -269,7 +269,6 @@ class TTSSpeakFrame(DataFrame):
@dataclass @dataclass
class TransportMessageFrame(DataFrame): class TransportMessageFrame(DataFrame):
message: Any message: Any
urgent: bool = False
def __str__(self): def __str__(self):
return f"{self.name}(message: {self.message})" return f"{self.name}(message: {self.message})"
@@ -405,6 +404,14 @@ class BotInterruptionFrame(SystemFrame):
pass pass
@dataclass
class TransportMessageUrgentFrame(SystemFrame):
message: Any
def __str__(self):
return f"{self.name}(message: {self.message})"
@dataclass @dataclass
class MetricsFrame(SystemFrame): class MetricsFrame(SystemFrame):
"""Emitted by processor that can compute metrics like latencies.""" """Emitted by processor that can compute metrics like latencies."""

View File

@@ -120,7 +120,7 @@ class ParallelPipeline(BasePipeline):
# If we get an EndFrame we stop our queue processing tasks and wait on # If we get an EndFrame we stop our queue processing tasks and wait on
# all the pipelines to finish. # all the pipelines to finish.
if isinstance(frame, CancelFrame) or isinstance(frame, EndFrame): if isinstance(frame, (CancelFrame, EndFrame)):
# Use None to indicate when queues should be done processing. # Use None to indicate when queues should be done processing.
await self._up_queue.put(None) await self._up_queue.put(None)
await self._down_queue.put(None) await self._down_queue.put(None)

View File

@@ -175,7 +175,7 @@ class PipelineTask:
await self._source.process_frame(frame, FrameDirection.DOWNSTREAM) await self._source.process_frame(frame, FrameDirection.DOWNSTREAM)
if isinstance(frame, EndFrame): if isinstance(frame, EndFrame):
await self._wait_for_endframe() await self._wait_for_endframe()
running = not (isinstance(frame, StopTaskFrame) or isinstance(frame, EndFrame)) running = not isinstance(frame, (StopTaskFrame, EndFrame))
should_cleanup = not isinstance(frame, StopTaskFrame) should_cleanup = not isinstance(frame, StopTaskFrame)
self._push_queue.task_done() self._push_queue.task_done()
except asyncio.CancelledError: except asyncio.CancelledError:

View File

@@ -6,10 +6,11 @@
import asyncio import asyncio
import base64 import base64
from typing import Any, Awaitable, Callable, Dict, List, Literal, Optional, Union
from pydantic import BaseModel, Field, PrivateAttr, ValidationError
from dataclasses import dataclass from dataclasses import dataclass
from typing import Any, Awaitable, Callable, Dict, List, Literal, Optional, Union
from loguru import logger
from pydantic import BaseModel, Field, PrivateAttr, ValidationError
from pipecat.frames.frames import ( from pipecat.frames.frames import (
BotInterruptionFrame, BotInterruptionFrame,
@@ -20,27 +21,28 @@ from pipecat.frames.frames import (
EndFrame, EndFrame,
ErrorFrame, ErrorFrame,
Frame, Frame,
FunctionCallResultFrame,
InterimTranscriptionFrame, InterimTranscriptionFrame,
LLMFullResponseEndFrame, LLMFullResponseEndFrame,
LLMFullResponseStartFrame, LLMFullResponseStartFrame,
OutputAudioRawFrame, OutputAudioRawFrame,
StartFrame, StartFrame,
SystemFrame, SystemFrame,
TTSStartedFrame,
TTSStoppedFrame,
TextFrame, TextFrame,
TranscriptionFrame, TranscriptionFrame,
TransportMessageFrame, TransportMessageFrame,
TransportMessageUrgentFrame,
TTSStartedFrame,
TTSStoppedFrame,
UserStartedSpeakingFrame, UserStartedSpeakingFrame,
FunctionCallResultFrame,
UserStoppedSpeakingFrame, UserStoppedSpeakingFrame,
) )
from pipecat.processors.aggregators.openai_llm_context import OpenAILLMContext from pipecat.processors.aggregators.openai_llm_context import (
OpenAILLMContext,
OpenAILLMContextFrame,
)
from pipecat.processors.frame_processor import FrameDirection, FrameProcessor from pipecat.processors.frame_processor import FrameDirection, FrameProcessor
from loguru import logger
RTVI_PROTOCOL_VERSION = "0.2" RTVI_PROTOCOL_VERSION = "0.2"
ActionResult = Union[bool, int, float, str, list, dict] ActionResult = Union[bool, int, float, str, list, dict]
@@ -291,22 +293,12 @@ class RTVIAudioMessageData(BaseModel):
num_channels: int num_channels: int
class RTVIBotAudioMessage(BaseModel): class RTVIBotTTSAudioMessage(BaseModel):
label: Literal["rtvi-ai"] = "rtvi-ai" label: Literal["rtvi-ai"] = "rtvi-ai"
type: Literal["bot-audio"] = "bot-audio" type: Literal["bot-tts-audio"] = "bot-tts-audio"
data: RTVIAudioMessageData data: RTVIAudioMessageData
class RTVIBotTranscriptionMessageData(BaseModel):
text: str
class RTVIBotTranscriptionMessage(BaseModel):
label: Literal["rtvi-ai"] = "rtvi-ai"
type: Literal["bot-transcription"] = "bot-transcription"
data: RTVIBotTranscriptionMessageData
class RTVIUserTranscriptionMessageData(BaseModel): class RTVIUserTranscriptionMessageData(BaseModel):
text: str text: str
user_id: str user_id: str
@@ -320,6 +312,12 @@ class RTVIUserTranscriptionMessage(BaseModel):
data: RTVIUserTranscriptionMessageData data: RTVIUserTranscriptionMessageData
class RTVIUserLLMTextMessage(BaseModel):
label: Literal["rtvi-ai"] = "rtvi-ai"
type: Literal["user-llm-text"] = "user-llm-text"
data: RTVITextMessageData
class RTVIUserStartedSpeakingMessage(BaseModel): class RTVIUserStartedSpeakingMessage(BaseModel):
label: Literal["rtvi-ai"] = "rtvi-ai" label: Literal["rtvi-ai"] = "rtvi-ai"
type: Literal["user-started-speaking"] = "user-started-speaking" type: Literal["user-started-speaking"] = "user-started-speaking"
@@ -350,9 +348,11 @@ class RTVIFrameProcessor(FrameProcessor):
self._direction = direction self._direction = direction
async def _push_transport_message(self, model: BaseModel, exclude_none: bool = True): async def _push_transport_message(self, model: BaseModel, exclude_none: bool = True):
frame = TransportMessageFrame( frame = TransportMessageFrame(message=model.model_dump(exclude_none=exclude_none))
message=model.model_dump(exclude_none=exclude_none), urgent=True await self.push_frame(frame, self._direction)
)
async def _push_transport_message_urgent(self, model: BaseModel, exclude_none: bool = True):
frame = TransportMessageUrgentFrame(message=model.model_dump(exclude_none=exclude_none))
await self.push_frame(frame, self._direction) await self.push_frame(frame, self._direction)
@@ -378,7 +378,7 @@ class RTVISpeakingProcessor(RTVIFrameProcessor):
message = RTVIUserStoppedSpeakingMessage() message = RTVIUserStoppedSpeakingMessage()
if message: if message:
await self._push_transport_message(message) await self._push_transport_message_urgent(message)
async def _handle_bot_speaking(self, frame: Frame): async def _handle_bot_speaking(self, frame: Frame):
message = None message = None
@@ -388,7 +388,7 @@ class RTVISpeakingProcessor(RTVIFrameProcessor):
message = RTVIBotStoppedSpeakingMessage() message = RTVIBotStoppedSpeakingMessage()
if message: if message:
await self._push_transport_message(message) await self._push_transport_message_urgent(message)
class RTVIUserTranscriptionProcessor(RTVIFrameProcessor): class RTVIUserTranscriptionProcessor(RTVIFrameProcessor):
@@ -419,7 +419,36 @@ class RTVIUserTranscriptionProcessor(RTVIFrameProcessor):
) )
if message: if message:
await self._push_transport_message(message) await self._push_transport_message_urgent(message)
class RTVIUserLLMTextProcessor(RTVIFrameProcessor):
def __init__(self, **kwargs):
super().__init__(**kwargs)
async def process_frame(self, frame: Frame, direction: FrameDirection):
await super().process_frame(frame, direction)
await self.push_frame(frame, direction)
if isinstance(frame, OpenAILLMContextFrame):
await self._handle_context(frame)
async def _handle_context(self, frame: OpenAILLMContextFrame):
messages = frame.context.messages
if len(messages) > 0:
message = messages[-1]
if message["role"] == "user":
content = message["content"]
if isinstance(content, list):
print("LIST")
text = " ".join(item["text"] for item in content if "text" in item)
else:
print("STRING")
text = content
rtvi_message = RTVIUserLLMTextMessage(data=RTVITextMessageData(text=text))
await self._push_transport_message_urgent(rtvi_message)
class RTVIBotLLMProcessor(RTVIFrameProcessor): class RTVIBotLLMProcessor(RTVIFrameProcessor):
@@ -432,9 +461,9 @@ class RTVIBotLLMProcessor(RTVIFrameProcessor):
await self.push_frame(frame, direction) await self.push_frame(frame, direction)
if isinstance(frame, LLMFullResponseStartFrame): if isinstance(frame, LLMFullResponseStartFrame):
await self._push_transport_message(RTVIBotLLMStartedMessage()) await self._push_transport_message_urgent(RTVIBotLLMStartedMessage())
elif isinstance(frame, LLMFullResponseEndFrame): elif isinstance(frame, LLMFullResponseEndFrame):
await self._push_transport_message(RTVIBotLLMStoppedMessage()) await self._push_transport_message_urgent(RTVIBotLLMStoppedMessage())
class RTVIBotTTSProcessor(RTVIFrameProcessor): class RTVIBotTTSProcessor(RTVIFrameProcessor):
@@ -447,9 +476,9 @@ class RTVIBotTTSProcessor(RTVIFrameProcessor):
await self.push_frame(frame, direction) await self.push_frame(frame, direction)
if isinstance(frame, TTSStartedFrame): if isinstance(frame, TTSStartedFrame):
await self._push_transport_message(RTVIBotTTSStartedMessage()) await self._push_transport_message_urgent(RTVIBotTTSStartedMessage())
elif isinstance(frame, TTSStoppedFrame): elif isinstance(frame, TTSStoppedFrame):
await self._push_transport_message(RTVIBotTTSStoppedMessage()) await self._push_transport_message_urgent(RTVIBotTTSStoppedMessage())
class RTVIBotLLMTextProcessor(RTVIFrameProcessor): class RTVIBotLLMTextProcessor(RTVIFrameProcessor):
@@ -466,7 +495,7 @@ class RTVIBotLLMTextProcessor(RTVIFrameProcessor):
async def _handle_text(self, frame: TextFrame): async def _handle_text(self, frame: TextFrame):
message = RTVIBotLLMTextMessage(data=RTVITextMessageData(text=frame.text)) message = RTVIBotLLMTextMessage(data=RTVITextMessageData(text=frame.text))
await self._push_transport_message(message) await self._push_transport_message_urgent(message)
class RTVIBotTTSTextProcessor(RTVIFrameProcessor): class RTVIBotTTSTextProcessor(RTVIFrameProcessor):
@@ -483,10 +512,10 @@ class RTVIBotTTSTextProcessor(RTVIFrameProcessor):
async def _handle_text(self, frame: TextFrame): async def _handle_text(self, frame: TextFrame):
message = RTVIBotTTSTextMessage(data=RTVITextMessageData(text=frame.text)) message = RTVIBotTTSTextMessage(data=RTVITextMessageData(text=frame.text))
await self._push_transport_message(message) await self._push_transport_message_urgent(message)
class RTVIBotAudioProcessor(RTVIFrameProcessor): class RTVIBotTTSAudioProcessor(RTVIFrameProcessor):
def __init__(self, **kwargs): def __init__(self, **kwargs):
super().__init__(**kwargs) super().__init__(**kwargs)
@@ -500,7 +529,7 @@ class RTVIBotAudioProcessor(RTVIFrameProcessor):
async def _handle_audio(self, frame: OutputAudioRawFrame): async def _handle_audio(self, frame: OutputAudioRawFrame):
encoded = base64.b64encode(frame.audio).decode("utf-8") encoded = base64.b64encode(frame.audio).decode("utf-8")
message = RTVIBotAudioMessage( message = RTVIBotTTSAudioMessage(
data=RTVIAudioMessageData( data=RTVIAudioMessageData(
audio=encoded, sample_rate=frame.sample_rate, num_channels=frame.num_channels audio=encoded, sample_rate=frame.sample_rate, num_channels=frame.num_channels
) )
@@ -647,9 +676,7 @@ class RTVIProcessor(FrameProcessor):
self._message_task = None self._message_task = None
async def _push_transport_message(self, model: BaseModel, exclude_none: bool = True): async def _push_transport_message(self, model: BaseModel, exclude_none: bool = True):
frame = TransportMessageFrame( frame = TransportMessageUrgentFrame(message=model.model_dump(exclude_none=exclude_none))
message=model.model_dump(exclude_none=exclude_none), urgent=True
)
await self.push_frame(frame) await self.push_frame(frame)
async def _action_task_handler(self): async def _action_task_handler(self):

View File

@@ -37,6 +37,7 @@ from pipecat.processors.frame_processor import FrameDirection, FrameProcessor
from pipecat.transcriptions.language import Language from pipecat.transcriptions.language import Language
from pipecat.utils.audio import calculate_audio_volume from pipecat.utils.audio import calculate_audio_volume
from pipecat.utils.string import match_endofsentence from pipecat.utils.string import match_endofsentence
from pipecat.utils.text.base_text_filter import BaseTextFilter
from pipecat.utils.time import seconds_to_nanoseconds from pipecat.utils.time import seconds_to_nanoseconds
from pipecat.utils.utils import exp_smoothing from pipecat.utils.utils import exp_smoothing
@@ -172,6 +173,7 @@ class TTSService(AIService):
stop_frame_timeout_s: float = 1.0, stop_frame_timeout_s: float = 1.0,
# TTS output sample rate # TTS output sample rate
sample_rate: int = 16000, sample_rate: int = 16000,
text_filter: Optional[BaseTextFilter] = None,
**kwargs, **kwargs,
): ):
super().__init__(**kwargs) super().__init__(**kwargs)
@@ -182,6 +184,7 @@ class TTSService(AIService):
self._sample_rate: int = sample_rate self._sample_rate: int = sample_rate
self._voice_id: str = "" self._voice_id: str = ""
self._settings: Dict[str, Any] = {} self._settings: Dict[str, Any] = {}
self._text_filter: Optional[BaseTextFilter] = text_filter
self._stop_frame_task: Optional[asyncio.Task] = None self._stop_frame_task: Optional[asyncio.Task] = None
self._stop_frame_queue: asyncio.Queue = asyncio.Queue() self._stop_frame_queue: asyncio.Queue = asyncio.Queue()
@@ -204,6 +207,9 @@ class TTSService(AIService):
async def flush_audio(self): async def flush_audio(self):
pass pass
def language_to_service_language(self, language: Language) -> str | None:
return Language(language)
# Converts the text to audio. # Converts the text to audio.
@abstractmethod @abstractmethod
async def run_tts(self, text: str) -> AsyncGenerator[Frame, None]: async def run_tts(self, text: str) -> AsyncGenerator[Frame, None]:
@@ -234,11 +240,13 @@ class TTSService(AIService):
logger.debug(f"Updating TTS setting {key} to: [{value}]") logger.debug(f"Updating TTS setting {key} to: [{value}]")
self._settings[key] = value self._settings[key] = value
if key == "language": if key == "language":
self._settings[key] = Language(value) self._settings[key] = self.language_to_service_language(value)
elif key == "model": elif key == "model":
self.set_model_name(value) self.set_model_name(value)
elif key == "voice": elif key == "voice":
self.set_voice(value) self.set_voice(value)
elif key == "text_filter" and self._text_filter:
self._text_filter.update_settings(value)
else: else:
logger.warning(f"Unknown setting for TTS service: {key}") logger.warning(f"Unknown setting for TTS service: {key}")
@@ -256,7 +264,7 @@ class TTSService(AIService):
await self._process_text_frame(frame) await self._process_text_frame(frame)
elif isinstance(frame, StartInterruptionFrame): elif isinstance(frame, StartInterruptionFrame):
await self._handle_interruption(frame, direction) await self._handle_interruption(frame, direction)
elif isinstance(frame, LLMFullResponseEndFrame) or isinstance(frame, EndFrame): elif isinstance(frame, (LLMFullResponseEndFrame, EndFrame)):
sentence = self._current_sentence sentence = self._current_sentence
self._current_sentence = "" self._current_sentence = ""
await self._push_tts_frames(sentence) await self._push_tts_frames(sentence)
@@ -309,6 +317,8 @@ class TTSService(AIService):
return return
await self.start_processing_metrics() await self.start_processing_metrics()
if self._text_filter:
text = self._text_filter.filter(text)
await self.process_generator(self.run_tts(text)) await self.process_generator(self.run_tts(text))
await self.stop_processing_metrics() await self.stop_processing_metrics()
if self._push_text_frames: if self._push_text_frames:
@@ -366,7 +376,7 @@ class WordTTSService(TTSService):
async def process_frame(self, frame: Frame, direction: FrameDirection): async def process_frame(self, frame: Frame, direction: FrameDirection):
await super().process_frame(frame, direction) await super().process_frame(frame, direction)
if isinstance(frame, LLMFullResponseEndFrame) or isinstance(frame, EndFrame): if isinstance(frame, (LLMFullResponseEndFrame, EndFrame)):
await self.flush_audio() await self.flush_audio()
async def _handle_interruption(self, frame: StartInterruptionFrame, direction: FrameDirection): async def _handle_interruption(self, frame: StartInterruptionFrame, direction: FrameDirection):
@@ -407,6 +417,10 @@ class STTService(AIService):
async def set_model(self, model: str): async def set_model(self, model: str):
self.set_model_name(model) self.set_model_name(model)
@abstractmethod
async def set_language(self, language: Language):
pass
@abstractmethod @abstractmethod
async def run_stt(self, audio: bytes) -> AsyncGenerator[Frame, None]: async def run_stt(self, audio: bytes) -> AsyncGenerator[Frame, None]:
"""Returns transcript as a string""" """Returns transcript as a string"""
@@ -419,7 +433,7 @@ class STTService(AIService):
logger.debug(f"Updating STT setting {key} to: [{value}]") logger.debug(f"Updating STT setting {key} to: [{value}]")
self._settings[key] = value self._settings[key] = value
if key == "language": if key == "language":
self._settings[key] = Language(value) await self.set_language(value)
elif key == "model": elif key == "model":
self.set_model_name(value) self.set_model_name(value)
else: else:

View File

@@ -30,67 +30,6 @@ except ModuleNotFoundError as e:
raise Exception(f"Missing module: {e}") raise Exception(f"Missing module: {e}")
def language_to_aws_language(language: Language) -> str | None:
match language:
case Language.CA:
return "ca-ES"
case Language.ZH:
return "cmn-CN"
case Language.DA:
return "da-DK"
case Language.NL:
return "nl-NL"
case Language.NL_BE:
return "nl-BE"
case Language.EN:
return "en-US"
case Language.EN_US:
return "en-US"
case Language.EN_AU:
return "en-AU"
case Language.EN_GB:
return "en-GB"
case Language.EN_NZ:
return "en-NZ"
case Language.EN_IN:
return "en-IN"
case Language.FI:
return "fi-FI"
case Language.FR:
return "fr-FR"
case Language.FR_CA:
return "fr-CA"
case Language.DE:
return "de-DE"
case Language.HI:
return "hi-IN"
case Language.IT:
return "it-IT"
case Language.JA:
return "ja-JP"
case Language.KO:
return "ko-KR"
case Language.NO:
return "nb-NO"
case Language.PL:
return "pl-PL"
case Language.PT:
return "pt-PT"
case Language.PT_BR:
return "pt-BR"
case Language.RO:
return "ro-RO"
case Language.RU:
return "ru-RU"
case Language.ES:
return "es-ES"
case Language.SV:
return "sv-SE"
case Language.TR:
return "tr-TR"
return None
class AWSTTSService(TTSService): class AWSTTSService(TTSService):
class InputParams(BaseModel): class InputParams(BaseModel):
engine: Optional[str] = None engine: Optional[str] = None
@@ -121,7 +60,9 @@ class AWSTTSService(TTSService):
self._settings = { self._settings = {
"sample_rate": sample_rate, "sample_rate": sample_rate,
"engine": params.engine, "engine": params.engine,
"language": params.language if params.language else Language.EN, "language": self.language_to_service_language(params.language)
if params.language
else Language.EN,
"pitch": params.pitch, "pitch": params.pitch,
"rate": params.rate, "rate": params.rate,
"volume": params.volume, "volume": params.volume,
@@ -132,10 +73,68 @@ class AWSTTSService(TTSService):
def can_generate_metrics(self) -> bool: def can_generate_metrics(self) -> bool:
return True return True
def language_to_service_language(self, language: Language) -> str | None:
match language:
case Language.CA:
return "ca-ES"
case Language.ZH:
return "cmn-CN"
case Language.DA:
return "da-DK"
case Language.NL:
return "nl-NL"
case Language.NL_BE:
return "nl-BE"
case Language.EN | Language.EN_US:
return "en-US"
case Language.EN_AU:
return "en-AU"
case Language.EN_GB:
return "en-GB"
case Language.EN_NZ:
return "en-NZ"
case Language.EN_IN:
return "en-IN"
case Language.FI:
return "fi-FI"
case Language.FR:
return "fr-FR"
case Language.FR_CA:
return "fr-CA"
case Language.DE:
return "de-DE"
case Language.HI:
return "hi-IN"
case Language.IT:
return "it-IT"
case Language.JA:
return "ja-JP"
case Language.KO:
return "ko-KR"
case Language.NO:
return "nb-NO"
case Language.PL:
return "pl-PL"
case Language.PT:
return "pt-PT"
case Language.PT_BR:
return "pt-BR"
case Language.RO:
return "ro-RO"
case Language.RU:
return "ru-RU"
case Language.ES:
return "es-ES"
case Language.SV:
return "sv-SE"
case Language.TR:
return "tr-TR"
return None
def _construct_ssml(self, text: str) -> str: def _construct_ssml(self, text: str) -> str:
ssml = "<speak>" ssml = "<speak>"
language = language_to_aws_language(self._settings["language"]) language = self._settings["language"]
ssml += f"<lang xml:lang='{language}'>" ssml += f"<lang xml:lang='{language}'>"
prosody_attrs = [] prosody_attrs = []

View File

@@ -27,7 +27,6 @@ from pipecat.frames.frames import (
) )
from pipecat.services.ai_services import ImageGenService, STTService, TTSService from pipecat.services.ai_services import ImageGenService, STTService, TTSService
from pipecat.services.openai import BaseOpenAILLMService from pipecat.services.openai import BaseOpenAILLMService
from pipecat.transcriptions import language
from pipecat.transcriptions.language import Language from pipecat.transcriptions.language import Language
from pipecat.utils.time import time_now_iso8601 from pipecat.utils.time import time_now_iso8601
@@ -72,101 +71,10 @@ class AzureLLMService(BaseOpenAILLMService):
) )
def language_to_azure_language(language: Language) -> str | None:
match language:
case Language.BG:
return "bg-BG"
case Language.CA:
return "ca-ES"
case Language.ZH:
return "zh-CN"
case Language.ZH_TW:
return "zh-TW"
case Language.CS:
return "cs-CZ"
case Language.DA:
return "da-DK"
case Language.NL:
return "nl-NL"
case Language.EN:
return "en-US"
case Language.EN_US:
return "en-US"
case Language.EN_AU:
return "en-AU"
case Language.EN_GB:
return "en-GB"
case Language.EN_NZ:
return "en-NZ"
case Language.EN_IN:
return "en-IN"
case Language.ET:
return "et-EE"
case Language.FI:
return "fi-FI"
case Language.NL_BE:
return "nl-BE"
case Language.FR:
return "fr-FR"
case Language.FR_CA:
return "fr-CA"
case Language.DE:
return "de-DE"
case Language.DE_CH:
return "de-CH"
case Language.EL:
return "el-GR"
case Language.HI:
return "hi-IN"
case Language.HU:
return "hu-HU"
case Language.ID:
return "id-ID"
case Language.IT:
return "it-IT"
case Language.JA:
return "ja-JP"
case Language.KO:
return "ko-KR"
case Language.LV:
return "lv-LV"
case Language.LT:
return "lt-LT"
case Language.MS:
return "ms-MY"
case Language.NO:
return "nb-NO"
case Language.PL:
return "pl-PL"
case Language.PT:
return "pt-PT"
case Language.PT_BR:
return "pt-BR"
case Language.RO:
return "ro-RO"
case Language.RU:
return "ru-RU"
case Language.SK:
return "sk-SK"
case Language.ES:
return "es-ES"
case Language.SV:
return "sv-SE"
case Language.TH:
return "th-TH"
case Language.TR:
return "tr-TR"
case Language.UK:
return "uk-UA"
case Language.VI:
return "vi-VN"
return None
class AzureTTSService(TTSService): class AzureTTSService(TTSService):
class InputParams(BaseModel): class InputParams(BaseModel):
emphasis: Optional[str] = None emphasis: Optional[str] = None
language: Optional[Language] = Language.EN language: Optional[Language] = Language.EN_US
pitch: Optional[str] = None pitch: Optional[str] = None
rate: Optional[str] = "1.05" rate: Optional[str] = "1.05"
role: Optional[str] = None role: Optional[str] = None
@@ -192,7 +100,9 @@ class AzureTTSService(TTSService):
self._settings = { self._settings = {
"sample_rate": sample_rate, "sample_rate": sample_rate,
"emphasis": params.emphasis, "emphasis": params.emphasis,
"language": params.language if params.language else Language.EN, "language": self.language_to_service_language(params.language)
if params.language
else Language.EN_US,
"pitch": params.pitch, "pitch": params.pitch,
"rate": params.rate, "rate": params.rate,
"role": params.role, "role": params.role,
@@ -206,8 +116,96 @@ class AzureTTSService(TTSService):
def can_generate_metrics(self) -> bool: def can_generate_metrics(self) -> bool:
return True return True
def language_to_service_language(self, language: Language) -> str | None:
match language:
case Language.BG:
return "bg-BG"
case Language.CA:
return "ca-ES"
case Language.ZH:
return "zh-CN"
case Language.ZH_TW:
return "zh-TW"
case Language.CS:
return "cs-CZ"
case Language.DA:
return "da-DK"
case Language.NL:
return "nl-NL"
case Language.EN | Language.EN_US:
return "en-US"
case Language.EN_AU:
return "en-AU"
case Language.EN_GB:
return "en-GB"
case Language.EN_NZ:
return "en-NZ"
case Language.EN_IN:
return "en-IN"
case Language.ET:
return "et-EE"
case Language.FI:
return "fi-FI"
case Language.NL_BE:
return "nl-BE"
case Language.FR:
return "fr-FR"
case Language.FR_CA:
return "fr-CA"
case Language.DE:
return "de-DE"
case Language.DE_CH:
return "de-CH"
case Language.EL:
return "el-GR"
case Language.HI:
return "hi-IN"
case Language.HU:
return "hu-HU"
case Language.ID:
return "id-ID"
case Language.IT:
return "it-IT"
case Language.JA:
return "ja-JP"
case Language.KO:
return "ko-KR"
case Language.LV:
return "lv-LV"
case Language.LT:
return "lt-LT"
case Language.MS:
return "ms-MY"
case Language.NO:
return "nb-NO"
case Language.PL:
return "pl-PL"
case Language.PT:
return "pt-PT"
case Language.PT_BR:
return "pt-BR"
case Language.RO:
return "ro-RO"
case Language.RU:
return "ru-RU"
case Language.SK:
return "sk-SK"
case Language.ES:
return "es-ES"
case Language.SV:
return "sv-SE"
case Language.TH:
return "th-TH"
case Language.TR:
return "tr-TR"
case Language.UK:
return "uk-UA"
case Language.VI:
return "vi-VN"
return None
def _construct_ssml(self, text: str) -> str: def _construct_ssml(self, text: str) -> str:
language = language_to_azure_language(self._settings["language"]) language = self._settings["language"]
ssml = ( ssml = (
f"<speak version='1.0' xml:lang='{language}' " f"<speak version='1.0' xml:lang='{language}' "
"xmlns='http://www.w3.org/2001/10/synthesis' " "xmlns='http://www.w3.org/2001/10/synthesis' "
@@ -284,7 +282,7 @@ class AzureSTTService(STTService):
*, *,
api_key: str, api_key: str,
region: str, region: str,
language="en-US", language=Language.EN_US,
sample_rate=16000, sample_rate=16000,
channels=1, channels=1,
**kwargs, **kwargs,

View File

@@ -45,17 +45,24 @@ def language_to_cartesia_language(language: Language) -> str | None:
match language: match language:
case Language.DE: case Language.DE:
return "de" return "de"
case Language.EN: case (
Language.EN
| Language.EN_US
| Language.EN_GB
| Language.EN_AU
| Language.EN_NZ
| Language.EN_IN
):
return "en" return "en"
case Language.ES: case Language.ES:
return "es" return "es"
case Language.FR: case Language.FR | Language.FR_CA:
return "fr" return "fr"
case Language.JA: case Language.JA:
return "ja" return "ja"
case Language.PT: case Language.PT | Language.PT_BR:
return "pt" return "pt"
case Language.ZH: case Language.ZH | Language.ZH_TW:
return "zh" return "zh"
return None return None
@@ -106,7 +113,9 @@ class CartesiaTTSService(WordTTSService):
"encoding": params.encoding, "encoding": params.encoding,
"sample_rate": params.sample_rate, "sample_rate": params.sample_rate,
}, },
"language": params.language if params.language else Language.EN, "language": self.language_to_service_language(params.language)
if params.language
else Language.EN,
"speed": params.speed, "speed": params.speed,
"emotion": params.emotion, "emotion": params.emotion,
} }
@@ -125,6 +134,9 @@ class CartesiaTTSService(WordTTSService):
await super().set_model(model) await super().set_model(model)
logger.debug(f"Switching TTS model to: [{model}]") logger.debug(f"Switching TTS model to: [{model}]")
def language_to_service_language(self, language: Language) -> str | None:
return language_to_cartesia_language(language)
def _build_msg( def _build_msg(
self, text: str = "", continue_transcript: bool = True, add_timestamps: bool = True self, text: str = "", continue_transcript: bool = True, add_timestamps: bool = True
): ):
@@ -146,7 +158,7 @@ class CartesiaTTSService(WordTTSService):
"model_id": self.model_name, "model_id": self.model_name,
"voice": voice_config, "voice": voice_config,
"output_format": self._settings["output_format"], "output_format": self._settings["output_format"],
"language": language_to_cartesia_language(self._settings["language"]), "language": self._settings["language"],
"add_timestamps": add_timestamps, "add_timestamps": add_timestamps,
} }
return json.dumps(msg) return json.dumps(msg)
@@ -303,7 +315,9 @@ class CartesiaHttpTTSService(TTSService):
"encoding": params.encoding, "encoding": params.encoding,
"sample_rate": params.sample_rate, "sample_rate": params.sample_rate,
}, },
"language": params.language if params.language else Language.EN, "language": self.language_to_service_language(params.language)
if params.language
else Language.EN,
"speed": params.speed, "speed": params.speed,
"emotion": params.emotion, "emotion": params.emotion,
} }
@@ -315,6 +329,9 @@ class CartesiaHttpTTSService(TTSService):
def can_generate_metrics(self) -> bool: def can_generate_metrics(self) -> bool:
return True return True
def language_to_service_language(self, language: Language) -> str | None:
return language_to_cartesia_language(language)
async def stop(self, frame: EndFrame): async def stop(self, frame: EndFrame):
await super().stop(frame) await super().stop(frame)
await self._client.close() await self._client.close()
@@ -343,7 +360,7 @@ class CartesiaHttpTTSService(TTSService):
transcript=text, transcript=text,
voice_id=self._voice_id, voice_id=self._voice_id,
output_format=self._settings["output_format"], output_format=self._settings["output_format"],
language=language_to_cartesia_language(self._settings["language"]), language=self._settings["language"],
stream=False, stream=False,
_experimental_voice_controls=voice_controls, _experimental_voice_controls=voice_controls,
) )

View File

@@ -158,6 +158,12 @@ class DeepgramSTTService(STTService):
await self._disconnect() await self._disconnect()
await self._connect() await self._connect()
async def set_language(self, language: Language):
logger.debug(f"Switching STT language to: [{language}]")
self._settings["language"] = language
await self._disconnect()
await self._connect()
async def start(self, frame: StartFrame): async def start(self, frame: StartFrame):
await super().start(frame) await super().start(frame)
await self._connect() await self._connect()

View File

@@ -50,76 +50,6 @@ def sample_rate_from_output_format(output_format: str) -> int:
return 16000 return 16000
def language_to_elevenlabs_language(language: Language) -> str | None:
match language:
case Language.BG:
return "bg"
case Language.ZH:
return "zh"
case Language.CS:
return "cs"
case Language.DA:
return "da"
case Language.NL:
return "nl"
case (
Language.EN
| Language.EN_US
| Language.EN_AU
| Language.EN_GB
| Language.EN_NZ
| Language.EN_IN
):
return "en"
case Language.FI:
return "fi"
case Language.FR | Language.FR_CA:
return "fr"
case Language.DE | Language.DE_CH:
return "de"
case Language.EL:
return "el"
case Language.HI:
return "hi"
case Language.HU:
return "hu"
case Language.ID:
return "id"
case Language.IT:
return "it"
case Language.JA:
return "ja"
case Language.KO:
return "ko"
case Language.MS:
return "ms"
case Language.NO:
return "no"
case Language.PL:
return "pl"
case Language.PT:
return "pt-PT"
case Language.PT_BR:
return "pt-BR"
case Language.RO:
return "ro"
case Language.RU:
return "ru"
case Language.SK:
return "sk"
case Language.ES:
return "es"
case Language.SV:
return "sv"
case Language.TR:
return "tr"
case Language.UK:
return "uk"
case Language.VI:
return "vi"
return None
def calculate_word_times( def calculate_word_times(
alignment_info: Mapping[str, Any], cumulative_time: float alignment_info: Mapping[str, Any], cumulative_time: float
) -> List[Tuple[str, float]]: ) -> List[Tuple[str, float]]:
@@ -198,7 +128,9 @@ class ElevenLabsTTSService(WordTTSService):
self._url = url self._url = url
self._settings = { self._settings = {
"sample_rate": sample_rate_from_output_format(params.output_format), "sample_rate": sample_rate_from_output_format(params.output_format),
"language": params.language if params.language else Language.EN, "language": self.language_to_service_language(params.language)
if params.language
else Language.EN,
"output_format": params.output_format, "output_format": params.output_format,
"optimize_streaming_latency": params.optimize_streaming_latency, "optimize_streaming_latency": params.optimize_streaming_latency,
"stability": params.stability, "stability": params.stability,
@@ -220,6 +152,75 @@ class ElevenLabsTTSService(WordTTSService):
def can_generate_metrics(self) -> bool: def can_generate_metrics(self) -> bool:
return True return True
def language_to_service_language(self, language: Language) -> str | None:
match language:
case Language.BG:
return "bg"
case Language.ZH:
return "zh"
case Language.CS:
return "cs"
case Language.DA:
return "da"
case Language.NL:
return "nl"
case (
Language.EN
| Language.EN_US
| Language.EN_AU
| Language.EN_GB
| Language.EN_NZ
| Language.EN_IN
):
return "en"
case Language.FI:
return "fi"
case Language.FR | Language.FR_CA:
return "fr"
case Language.DE | Language.DE_CH:
return "de"
case Language.EL:
return "el"
case Language.HI:
return "hi"
case Language.HU:
return "hu"
case Language.ID:
return "id"
case Language.IT:
return "it"
case Language.JA:
return "ja"
case Language.KO:
return "ko"
case Language.MS:
return "ms"
case Language.NO:
return "no"
case Language.PL:
return "pl"
case Language.PT:
return "pt-PT"
case Language.PT_BR:
return "pt-BR"
case Language.RO:
return "ro"
case Language.RU:
return "ru"
case Language.SK:
return "sk"
case Language.ES:
return "es"
case Language.SV:
return "sv"
case Language.TR:
return "tr"
case Language.UK:
return "uk"
case Language.VI:
return "vi"
return None
def _set_voice_settings(self): def _set_voice_settings(self):
voice_settings = {} voice_settings = {}
if ( if (
@@ -293,7 +294,7 @@ class ElevenLabsTTSService(WordTTSService):
url += f"&optimize_streaming_latency={self._settings['optimize_streaming_latency']}" url += f"&optimize_streaming_latency={self._settings['optimize_streaming_latency']}"
# Language can only be used with the 'eleven_turbo_v2_5' model # Language can only be used with the 'eleven_turbo_v2_5' model
language = language_to_elevenlabs_language(self._settings["language"]) language = self._settings["language"]
if model == "eleven_turbo_v2_5": if model == "eleven_turbo_v2_5":
url += f"&language_code={language}" url += f"&language_code={language}"
else: else:

View File

@@ -34,84 +34,6 @@ except ModuleNotFoundError as e:
raise Exception(f"Missing module: {e}") raise Exception(f"Missing module: {e}")
def language_to_gladia_language(language: Language) -> str | None:
match language:
case Language.BG:
return "bulgarian"
case Language.CA:
return "catalan"
case Language.ZH:
return "chinese"
case Language.CS:
return "czech"
case Language.DA:
return "danish"
case Language.NL:
return "dutch"
case (
Language.EN
| Language.EN_US
| Language.EN_AU
| Language.EN_GB
| Language.EN_NZ
| Language.EN_IN
):
return "english"
case Language.ET:
return "estonian"
case Language.FI:
return "finnish"
case Language.FR | Language.FR_CA:
return "french"
case Language.DE | Language.DE_CH:
return "german"
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.LV:
return "latvian"
case Language.LT:
return "lithuanian"
case Language.MS:
return "malay"
case Language.NO:
return "norwegian"
case Language.PL:
return "polish"
case Language.PT | Language.PT_BR:
return "portuguese"
case Language.RO:
return "romanian"
case Language.RU:
return "russian"
case Language.SK:
return "slovak"
case Language.ES:
return "spanish"
case Language.SV:
return "slovenian"
case Language.TH:
return "thai"
case Language.TR:
return "turkish"
case Language.UK:
return "ukrainian"
case Language.VI:
return "vietnamese"
return None
class GladiaSTTService(STTService): class GladiaSTTService(STTService):
class InputParams(BaseModel): class InputParams(BaseModel):
sample_rate: Optional[int] = 16000 sample_rate: Optional[int] = 16000
@@ -135,13 +57,92 @@ class GladiaSTTService(STTService):
self._url = url self._url = url
self._settings = { self._settings = {
"sample_rate": params.sample_rate, "sample_rate": params.sample_rate,
"language": params.language if params.language else Language.EN, "language": self.language_to_service_language(params.language)
if params.language
else Language.EN,
"transcription_hint": params.transcription_hint, "transcription_hint": params.transcription_hint,
"endpointing": params.endpointing, "endpointing": params.endpointing,
"prosody": params.prosody, "prosody": params.prosody,
} }
self._confidence = confidence self._confidence = confidence
def language_to_service_language(self, language: Language) -> str | None:
match language:
case Language.BG:
return "bulgarian"
case Language.CA:
return "catalan"
case Language.ZH:
return "chinese"
case Language.CS:
return "czech"
case Language.DA:
return "danish"
case Language.NL:
return "dutch"
case (
Language.EN
| Language.EN_US
| Language.EN_AU
| Language.EN_GB
| Language.EN_NZ
| Language.EN_IN
):
return "english"
case Language.ET:
return "estonian"
case Language.FI:
return "finnish"
case Language.FR | Language.FR_CA:
return "french"
case Language.DE | Language.DE_CH:
return "german"
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.LV:
return "latvian"
case Language.LT:
return "lithuanian"
case Language.MS:
return "malay"
case Language.NO:
return "norwegian"
case Language.PL:
return "polish"
case Language.PT | Language.PT_BR:
return "portuguese"
case Language.RO:
return "romanian"
case Language.RU:
return "russian"
case Language.SK:
return "slovak"
case Language.ES:
return "spanish"
case Language.SV:
return "slovenian"
case Language.TH:
return "thai"
case Language.TR:
return "turkish"
case Language.UK:
return "ukrainian"
case Language.VI:
return "vietnamese"
return None
async def start(self, frame: StartFrame): async def start(self, frame: StartFrame):
await super().start(frame) await super().start(frame)
self._websocket = await websockets.connect(self._url) self._websocket = await websockets.connect(self._url)
@@ -169,7 +170,7 @@ class GladiaSTTService(STTService):
"model_type": "fast", "model_type": "fast",
"language_behaviour": "manual", "language_behaviour": "manual",
"sample_rate": self._settings["sample_rate"], "sample_rate": self._settings["sample_rate"],
"language": language_to_gladia_language(self._settings["language"]), "language": self._settings["language"],
"transcription_hint": self._settings["transcription_hint"], "transcription_hint": self._settings["transcription_hint"],
"endpointing": self._settings["endpointing"], "endpointing": self._settings["endpointing"],
"prosody": self._settings["prosody"], "prosody": self._settings["prosody"],

View File

@@ -146,93 +146,6 @@ class GoogleLLMService(LLMService):
await self._process_context(context) await self._process_context(context)
def language_to_google_language(language: Language) -> str | None:
match language:
case Language.BG:
return "bg-BG"
case Language.CA:
return "ca-ES"
case Language.ZH:
return "cmn-CN"
case Language.ZH_TW:
return "cmn-TW"
case Language.CS:
return "cs-CZ"
case Language.DA:
return "da-DK"
case Language.NL:
return "nl-NL"
case Language.EN:
return "en-US"
case Language.EN_US:
return "en-US"
case Language.EN_AU:
return "en-AU"
case Language.EN_GB:
return "en-GB"
case Language.EN_IN:
return "en-IN"
case Language.ET:
return "et-EE"
case Language.FI:
return "fi-FI"
case Language.NL_BE:
return "nl-BE"
case Language.FR:
return "fr-FR"
case Language.FR_CA:
return "fr-CA"
case Language.DE:
return "de-DE"
case Language.EL:
return "el-GR"
case Language.HI:
return "hi-IN"
case Language.HU:
return "hu-HU"
case Language.ID:
return "id-ID"
case Language.IT:
return "it-IT"
case Language.JA:
return "ja-JP"
case Language.KO:
return "ko-KR"
case Language.LV:
return "lv-LV"
case Language.LT:
return "lt-LT"
case Language.MS:
return "ms-MY"
case Language.NO:
return "nb-NO"
case Language.PL:
return "pl-PL"
case Language.PT:
return "pt-PT"
case Language.PT_BR:
return "pt-BR"
case Language.RO:
return "ro-RO"
case Language.RU:
return "ru-RU"
case Language.SK:
return "sk-SK"
case Language.ES:
return "es-ES"
case Language.SV:
return "sv-SE"
case Language.TH:
return "th-TH"
case Language.TR:
return "tr-TR"
case Language.UK:
return "uk-UA"
case Language.VI:
return "vi-VN"
return None
class GoogleTTSService(TTSService): class GoogleTTSService(TTSService):
class InputParams(BaseModel): class InputParams(BaseModel):
pitch: Optional[str] = None pitch: Optional[str] = None
@@ -261,7 +174,9 @@ class GoogleTTSService(TTSService):
"rate": params.rate, "rate": params.rate,
"volume": params.volume, "volume": params.volume,
"emphasis": params.emphasis, "emphasis": params.emphasis,
"language": params.language if params.language else Language.EN, "language": self.language_to_service_language(params.language)
if params.language
else Language.EN,
"gender": params.gender, "gender": params.gender,
"google_style": params.google_style, "google_style": params.google_style,
} }
@@ -291,13 +206,97 @@ class GoogleTTSService(TTSService):
def can_generate_metrics(self) -> bool: def can_generate_metrics(self) -> bool:
return True return True
def language_to_service_language(self, language: Language) -> str | None:
match language:
case Language.BG:
return "bg-BG"
case Language.CA:
return "ca-ES"
case Language.ZH:
return "cmn-CN"
case Language.ZH_TW:
return "cmn-TW"
case Language.CS:
return "cs-CZ"
case Language.DA:
return "da-DK"
case Language.NL:
return "nl-NL"
case Language.EN | Language.EN_US:
return "en-US"
case Language.EN_AU:
return "en-AU"
case Language.EN_GB:
return "en-GB"
case Language.EN_IN:
return "en-IN"
case Language.ET:
return "et-EE"
case Language.FI:
return "fi-FI"
case Language.NL_BE:
return "nl-BE"
case Language.FR:
return "fr-FR"
case Language.FR_CA:
return "fr-CA"
case Language.DE:
return "de-DE"
case Language.EL:
return "el-GR"
case Language.HI:
return "hi-IN"
case Language.HU:
return "hu-HU"
case Language.ID:
return "id-ID"
case Language.IT:
return "it-IT"
case Language.JA:
return "ja-JP"
case Language.KO:
return "ko-KR"
case Language.LV:
return "lv-LV"
case Language.LT:
return "lt-LT"
case Language.MS:
return "ms-MY"
case Language.NO:
return "nb-NO"
case Language.PL:
return "pl-PL"
case Language.PT:
return "pt-PT"
case Language.PT_BR:
return "pt-BR"
case Language.RO:
return "ro-RO"
case Language.RU:
return "ru-RU"
case Language.SK:
return "sk-SK"
case Language.ES:
return "es-ES"
case Language.SV:
return "sv-SE"
case Language.TH:
return "th-TH"
case Language.TR:
return "tr-TR"
case Language.UK:
return "uk-UA"
case Language.VI:
return "vi-VN"
return None
def _construct_ssml(self, text: str) -> str: def _construct_ssml(self, text: str) -> str:
ssml = "<speak>" ssml = "<speak>"
# Voice tag # Voice tag
voice_attrs = [f"name='{self._voice_id}'"] voice_attrs = [f"name='{self._voice_id}'"]
language = language_to_google_language(self._settings["language"]) language = self._settings["language"]
voice_attrs.append(f"language='{language}'") voice_attrs.append(f"language='{language}'")
if self._settings["gender"]: if self._settings["gender"]:

View File

@@ -35,32 +35,6 @@ except ModuleNotFoundError as e:
raise Exception(f"Missing module: {e}") raise Exception(f"Missing module: {e}")
def language_to_lmnt_language(language: Language) -> str | None:
match language:
case Language.DE:
return "de"
case (
Language.EN
| Language.EN_US
| Language.EN_AU
| Language.EN_GB
| Language.EN_NZ
| Language.EN_IN
):
return "en"
case Language.ES:
return "es"
case Language.FR | Language.FR_CA:
return "fr"
case Language.PT | Language.PT_BR:
return "pt"
case Language.ZH | Language.ZH_TW:
return "zh"
case Language.KO:
return "ko"
return None
class LmntTTSService(TTSService): class LmntTTSService(TTSService):
def __init__( def __init__(
self, self,
@@ -82,7 +56,7 @@ class LmntTTSService(TTSService):
"encoding": "pcm_s16le", "encoding": "pcm_s16le",
"sample_rate": sample_rate, "sample_rate": sample_rate,
}, },
"language": language, "language": self.language_to_service_language(language),
} }
self.set_voice(voice_id) self.set_voice(voice_id)
@@ -97,6 +71,31 @@ class LmntTTSService(TTSService):
def can_generate_metrics(self) -> bool: def can_generate_metrics(self) -> bool:
return True return True
def language_to_service_language(self, language: Language) -> str | None:
match language:
case Language.DE:
return "de"
case (
Language.EN
| Language.EN_US
| Language.EN_AU
| Language.EN_GB
| Language.EN_NZ
| Language.EN_IN
):
return "en"
case Language.ES:
return "es"
case Language.FR | Language.FR_CA:
return "fr"
case Language.PT | Language.PT_BR:
return "pt"
case Language.ZH | Language.ZH_TW:
return "zh"
case Language.KO:
return "ko"
return None
async def start(self, frame: StartFrame): async def start(self, frame: StartFrame):
await super().start(frame) await super().start(frame)
await self._connect() await self._connect()
@@ -121,6 +120,7 @@ class LmntTTSService(TTSService):
self._voice_id, self._voice_id,
format="raw", format="raw",
sample_rate=self._settings["output_format"]["sample_rate"], sample_rate=self._settings["output_format"]["sample_rate"],
language=self._settings["language"],
) )
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())
except Exception as e: except Exception as e:

View File

@@ -37,50 +37,6 @@ except ModuleNotFoundError as e:
# https://github.com/coqui-ai/xtts-streaming-server # https://github.com/coqui-ai/xtts-streaming-server
def language_to_xtts_language(language: Language) -> str | None:
match language:
case Language.CS:
return "cs"
case Language.DE:
return "de"
case (
Language.EN
| Language.EN_US
| Language.EN_AU
| Language.EN_GB
| Language.EN_NZ
| Language.EN_IN
):
return "en"
case Language.ES:
return "es"
case Language.FR:
return "fr"
case Language.HI:
return "hi"
case Language.HU:
return "hu"
case Language.IT:
return "it"
case Language.JA:
return "ja"
case Language.KO:
return "ko"
case Language.NL:
return "nl"
case Language.PL:
return "pl"
case Language.PT | Language.PT_BR:
return "pt"
case Language.RU:
return "ru"
case Language.TR:
return "tr"
case Language.ZH:
return "zh-cn"
return None
class XTTSService(TTSService): class XTTSService(TTSService):
def __init__( def __init__(
self, self,
@@ -94,7 +50,7 @@ class XTTSService(TTSService):
super().__init__(**kwargs) super().__init__(**kwargs)
self._settings = { self._settings = {
"language": language, "language": self.language_to_service_language(language),
"base_url": base_url, "base_url": base_url,
} }
self.set_voice(voice_id) self.set_voice(voice_id)
@@ -104,6 +60,49 @@ class XTTSService(TTSService):
def can_generate_metrics(self) -> bool: def can_generate_metrics(self) -> bool:
return True return True
def language_to_service_language(self, language: Language) -> str | None:
match language:
case Language.CS:
return "cs"
case Language.DE:
return "de"
case (
Language.EN
| Language.EN_US
| Language.EN_AU
| Language.EN_GB
| Language.EN_NZ
| Language.EN_IN
):
return "en"
case Language.ES:
return "es"
case Language.FR:
return "fr"
case Language.HI:
return "hi"
case Language.HU:
return "hu"
case Language.IT:
return "it"
case Language.JA:
return "ja"
case Language.KO:
return "ko"
case Language.NL:
return "nl"
case Language.PL:
return "pl"
case Language.PT | Language.PT_BR:
return "pt"
case Language.RU:
return "ru"
case Language.TR:
return "tr"
case Language.ZH:
return "zh-cn"
return None
async def start(self, frame: StartFrame): async def start(self, frame: StartFrame):
await super().start(frame) await super().start(frame)
async with self._aiohttp_session.get(self._settings["base_url"] + "/studio_speakers") as r: async with self._aiohttp_session.get(self._settings["base_url"] + "/studio_speakers") as r:
@@ -131,11 +130,9 @@ class XTTSService(TTSService):
url = self._settings["base_url"] + "/tts_stream" url = self._settings["base_url"] + "/tts_stream"
language = language_to_xtts_language(self._settings["language"])
payload = { payload = {
"text": text.replace(".", "").replace("*", ""), "text": text.replace(".", "").replace("*", ""),
"language": language, "language": self._settings["language"],
"speaker_embedding": embeddings["speaker_embedding"], "speaker_embedding": embeddings["speaker_embedding"],
"gpt_cond_latent": embeddings["gpt_cond_latent"], "gpt_cond_latent": embeddings["gpt_cond_latent"],
"add_wav_header": False, "add_wav_header": False,

View File

@@ -33,6 +33,7 @@ from pipecat.frames.frames import (
TTSStoppedFrame, TTSStoppedFrame,
TextFrame, TextFrame,
TransportMessageFrame, TransportMessageFrame,
TransportMessageUrgentFrame,
) )
from pipecat.transports.base_transport import TransportParams from pipecat.transports.base_transport import TransportParams
@@ -148,7 +149,7 @@ class BaseOutputTransport(FrameProcessor):
await self._audio_out_task await self._audio_out_task
self._audio_out_task = None self._audio_out_task = None
async def send_message(self, frame: TransportMessageFrame): async def send_message(self, frame: TransportMessageFrame | TransportMessageUrgentFrame):
pass pass
async def send_metrics(self, frame: MetricsFrame): async def send_metrics(self, frame: MetricsFrame):
@@ -180,12 +181,14 @@ class BaseOutputTransport(FrameProcessor):
elif isinstance(frame, CancelFrame): elif isinstance(frame, CancelFrame):
await self.cancel(frame) await self.cancel(frame)
await self.push_frame(frame, direction) await self.push_frame(frame, direction)
elif isinstance(frame, StartInterruptionFrame) or isinstance(frame, StopInterruptionFrame): elif isinstance(frame, (StartInterruptionFrame, StopInterruptionFrame)):
await self.push_frame(frame, direction) await self.push_frame(frame, direction)
await self._handle_interruptions(frame) await self._handle_interruptions(frame)
elif isinstance(frame, MetricsFrame): elif isinstance(frame, MetricsFrame):
await self.push_frame(frame, direction) await self.push_frame(frame, direction)
await self.send_metrics(frame) await self.send_metrics(frame)
elif isinstance(frame, TransportMessageUrgentFrame):
await self.send_message(frame)
elif isinstance(frame, SystemFrame): elif isinstance(frame, SystemFrame):
await self.push_frame(frame, direction) await self.push_frame(frame, direction)
# Control frames. # Control frames.
@@ -196,10 +199,8 @@ class BaseOutputTransport(FrameProcessor):
# Other frames. # Other frames.
elif isinstance(frame, OutputAudioRawFrame): elif isinstance(frame, OutputAudioRawFrame):
await self._handle_audio(frame) await self._handle_audio(frame)
elif isinstance(frame, OutputImageRawFrame) or isinstance(frame, SpriteFrame): elif isinstance(frame, (OutputImageRawFrame, SpriteFrame)):
await self._handle_image(frame) await self._handle_image(frame)
elif isinstance(frame, TransportMessageFrame) and frame.urgent:
await self.send_message(frame)
# TODO(aleix): Images and audio should support presentation timestamps. # TODO(aleix): Images and audio should support presentation timestamps.
elif frame.pts: elif frame.pts:
await self._sink_clock_queue.put((frame.pts, frame.id, frame)) await self._sink_clock_queue.put((frame.pts, frame.id, frame))

View File

@@ -35,6 +35,7 @@ from pipecat.frames.frames import (
StartFrame, StartFrame,
TranscriptionFrame, TranscriptionFrame,
TransportMessageFrame, TransportMessageFrame,
TransportMessageUrgentFrame,
UserImageRawFrame, UserImageRawFrame,
UserImageRequestFrame, UserImageRequestFrame,
) )
@@ -70,6 +71,11 @@ class DailyTransportMessageFrame(TransportMessageFrame):
participant_id: str | None = None participant_id: str | None = None
@dataclass
class DailyTransportMessageUrgentFrame(TransportMessageUrgentFrame):
participant_id: str | None = None
class WebRTCVADAnalyzer(VADAnalyzer): class WebRTCVADAnalyzer(VADAnalyzer):
def __init__(self, *, sample_rate=16000, num_channels=1, params: VADParams = VADParams()): def __init__(self, *, sample_rate=16000, num_channels=1, params: VADParams = VADParams()):
super().__init__(sample_rate=sample_rate, num_channels=num_channels, params=params) super().__init__(sample_rate=sample_rate, num_channels=num_channels, params=params)
@@ -234,12 +240,12 @@ class DailyTransportClient(EventHandler):
def set_callbacks(self, callbacks: DailyCallbacks): def set_callbacks(self, callbacks: DailyCallbacks):
self._callbacks = callbacks self._callbacks = callbacks
async def send_message(self, frame: TransportMessageFrame): async def send_message(self, frame: TransportMessageFrame | TransportMessageUrgentFrame):
if not self._client: if not self._client:
return return
participant_id = None participant_id = None
if isinstance(frame, DailyTransportMessageFrame): if isinstance(frame, (DailyTransportMessageFrame, DailyTransportMessageUrgentFrame)):
participant_id = frame.participant_id participant_id = frame.participant_id
future = self._loop.create_future() future = self._loop.create_future()
@@ -736,7 +742,7 @@ class DailyOutputTransport(BaseOutputTransport):
await super().cleanup() await super().cleanup()
await self._client.cleanup() await self._client.cleanup()
async def send_message(self, frame: TransportMessageFrame): async def send_message(self, frame: TransportMessageFrame | TransportMessageUrgentFrame):
await self._client.send_message(frame) await self._client.send_message(frame)
async def send_metrics(self, frame: MetricsFrame): async def send_metrics(self, frame: MetricsFrame):

View File

@@ -22,6 +22,7 @@ from pipecat.frames.frames import (
MetricsFrame, MetricsFrame,
StartFrame, StartFrame,
TransportMessageFrame, TransportMessageFrame,
TransportMessageUrgentFrame,
) )
from pipecat.metrics.metrics import ( from pipecat.metrics.metrics import (
LLMUsageMetricsData, LLMUsageMetricsData,
@@ -51,6 +52,11 @@ class LiveKitTransportMessageFrame(TransportMessageFrame):
participant_id: str | None = None participant_id: str | None = None
@dataclass
class LiveKitTransportMessageUrgentFrame(TransportMessageUrgentFrame):
participant_id: str | None = None
class LiveKitParams(TransportParams): class LiveKitParams(TransportParams):
audio_out_sample_rate: int = 48000 audio_out_sample_rate: int = 48000
audio_out_channels: int = 1 audio_out_channels: int = 1
@@ -420,8 +426,8 @@ class LiveKitOutputTransport(BaseOutputTransport):
await super().cancel(frame) await super().cancel(frame)
await self._client.disconnect() await self._client.disconnect()
async def send_message(self, frame: TransportMessageFrame): async def send_message(self, frame: TransportMessageFrame | TransportMessageUrgentFrame):
if isinstance(frame, LiveKitTransportMessageFrame): if isinstance(frame, (LiveKitTransportMessageFrame, LiveKitTransportMessageUrgentFrame)):
await self._client.send_data(frame.message.encode(), frame.participant_id) await self._client.send_data(frame.message.encode(), frame.participant_id)
else: else:
await self._client.send_data(frame.message.encode()) await self._client.send_data(frame.message.encode())
@@ -596,6 +602,13 @@ class LiveKitTransport(BaseTransport):
frame = LiveKitTransportMessageFrame(message=message, participant_id=participant_id) frame = LiveKitTransportMessageFrame(message=message, participant_id=participant_id)
await self._output.send_message(frame) await self._output.send_message(frame)
async def send_message_urgent(self, message: str, participant_id: str | None = None):
if self._output:
frame = LiveKitTransportMessageUrgentFrame(
message=message, participant_id=participant_id
)
await self._output.send_message(frame)
async def cleanup(self): async def cleanup(self):
if self._input: if self._input:
await self._input.cleanup() await self._input.cleanup()

View File

View File

@@ -0,0 +1,18 @@
#
# Copyright (c) 2024, Daily
#
# SPDX-License-Identifier: BSD 2-Clause License
#
from abc import ABC, abstractmethod
from typing import Any, Mapping
class BaseTextFilter(ABC):
@abstractmethod
def update_settings(self, settings: Mapping[str, Any]):
pass
@abstractmethod
def filter(self, text: str) -> str:
pass

View File

@@ -0,0 +1,84 @@
#
# Copyright (c) 2024, Daily
#
# SPDX-License-Identifier: BSD 2-Clause License
#
import re
from typing import Any, Mapping
from markdown import Markdown
from pydantic import BaseModel
from pipecat.utils.text.base_text_filter import BaseTextFilter
class MarkdownTextFilter(BaseTextFilter):
"""Removes Markdown formatting from text in TextFrames.
Converts Markdown to plain text while preserving the overall structure,
including leading and trailing spaces. Handles special cases like
asterisks and table formatting.
"""
class InputParams(BaseModel):
enable_text_filter: bool = True
def __init__(self, params: InputParams = InputParams(), **kwargs):
super().__init__(**kwargs)
self._settings = params
def update_settings(self, settings: Mapping[str, Any]):
for key, value in settings.items():
if hasattr(self._settings, key):
setattr(self._settings, key, value)
def filter(self, text: str) -> str:
if self._settings.enable_text_filter:
# Replace newlines with spaces only when there's no text before or after
text = re.sub(r"^\s*\n", " ", text, flags=re.MULTILINE)
# Remove repeated sequences of 5 or more characters
text = re.sub(r"(\S)(\1{4,})", "", text)
# Preserve numbered list items with a unique marker, §NUM§
text = re.sub(r"^(\d+\.)\s", r"§NUM§\1 ", text)
# Preserve leading/trailing spaces with a unique marker, §
# Critical for word-by-word streaming in bot-tts-text
preserved_markdown = re.sub(
r"^( +)|\s+$", lambda m: "§" * len(m.group(0)), text, flags=re.MULTILINE
)
# Convert markdown to HTML
md = Markdown()
html = md.convert(preserved_markdown)
# Remove HTML tags
filtered_text = re.sub("<[^<]+?>", "", html)
# Replace HTML entities
filtered_text = filtered_text.replace("&nbsp;", " ")
filtered_text = filtered_text.replace("&lt;", "<")
filtered_text = filtered_text.replace("&gt;", ">")
filtered_text = filtered_text.replace("&amp;", "&")
# Remove double asterisks (consecutive without any exceptions)
filtered_text = re.sub(r"\*\*", "", filtered_text)
# Remove single asterisks at the start or end of words
filtered_text = re.sub(r"(^|\s)\*|\*($|\s)", r"\1\2", filtered_text)
# Remove Markdown table formatting
filtered_text = re.sub(r"\|", "", filtered_text)
filtered_text = re.sub(r"^\s*[-:]+\s*$", "", filtered_text, flags=re.MULTILINE)
# Restore numbered list items
filtered_text = filtered_text.replace("§NUM§", "")
# Restore leading and trailing spaces
filtered_text = re.sub("§", " ", filtered_text)
return filtered_text
else:
return text

View File

@@ -5,7 +5,7 @@ boto3~=1.35.27
daily-python~=0.11.0 daily-python~=0.11.0
deepgram-sdk~=3.5.0 deepgram-sdk~=3.5.0
fal-client~=0.4.1 fal-client~=0.4.1
fastapi~=0.112.1 fastapi~=0.115.0
faster-whisper~=1.0.3 faster-whisper~=1.0.3
google-cloud-texttospeech~=2.17.2 google-cloud-texttospeech~=2.17.2
google-generativeai~=0.7.2 google-generativeai~=0.7.2