Merge branch 'main' into recording
This commit is contained in:
26
CHANGELOG.md
26
CHANGELOG.md
@@ -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
|
||||||
|
|||||||
@@ -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" ]
|
||||||
|
|||||||
@@ -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."""
|
||||||
|
|||||||
@@ -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)
|
||||||
|
|||||||
@@ -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:
|
||||||
|
|||||||
@@ -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):
|
||||||
|
|||||||
@@ -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:
|
||||||
|
|||||||
@@ -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 = []
|
||||||
|
|||||||
@@ -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,
|
||||||
|
|||||||
@@ -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,
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -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()
|
||||||
|
|||||||
@@ -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:
|
||||||
|
|||||||
@@ -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"],
|
||||||
|
|||||||
@@ -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"]:
|
||||||
|
|||||||
@@ -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:
|
||||||
|
|||||||
@@ -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,
|
||||||
|
|||||||
@@ -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))
|
||||||
|
|||||||
@@ -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):
|
||||||
|
|||||||
@@ -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()
|
||||||
|
|||||||
0
src/pipecat/utils/text/__init__.py
Normal file
0
src/pipecat/utils/text/__init__.py
Normal file
18
src/pipecat/utils/text/base_text_filter.py
Normal file
18
src/pipecat/utils/text/base_text_filter.py
Normal 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
|
||||||
84
src/pipecat/utils/text/markdown_text_filter.py
Normal file
84
src/pipecat/utils/text/markdown_text_filter.py
Normal 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(" ", " ")
|
||||||
|
filtered_text = filtered_text.replace("<", "<")
|
||||||
|
filtered_text = filtered_text.replace(">", ">")
|
||||||
|
filtered_text = filtered_text.replace("&", "&")
|
||||||
|
|
||||||
|
# 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
|
||||||
@@ -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
|
||||||
|
|||||||
Reference in New Issue
Block a user