Adding support for new bot-output RTVI Message:
1. TTSTextFrames now include metadata about whether the text was spoken or not along with a type string to describe what the text represents: ex. "sentence", "word", "custom aggregation" 2. Expanded how aggregators work so that the aggregate method returns aggregated text along with the type of aggregation used to create it 3. Deprecated the RTVI bot-transcription event in lieu of... 4. Introduced support for a new bot-output event. This event is meant to be the one stop shop for communicating what the bot actually "says". It is based off TTSTextFrames to communicate both sentence by sentence (or whatever aggregation is used) as well as word by word. In addition, it will include LLMTextFrames, aggregated by sentence when tts is turned off (i.e. skip_tts is true). Resolves pipecat-ai/pipecat-client-web#158
This commit is contained in:
@@ -359,6 +359,9 @@ class LLMTextFrame(TextFrame):
|
|||||||
class TTSTextFrame(TextFrame):
|
class TTSTextFrame(TextFrame):
|
||||||
"""Text frame generated by Text-to-Speech services."""
|
"""Text frame generated by Text-to-Speech services."""
|
||||||
|
|
||||||
|
aggregated_by: Literal["sentence", "word"] | str
|
||||||
|
spoken: Optional[bool] = True # Whether this text has been spoken by TTS
|
||||||
|
|
||||||
pass
|
pass
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -704,6 +704,29 @@ class RTVITextMessageData(BaseModel):
|
|||||||
text: str
|
text: str
|
||||||
|
|
||||||
|
|
||||||
|
class RTVIBotOutputMessageData(RTVITextMessageData):
|
||||||
|
"""Data for bot output RTVI messages.
|
||||||
|
|
||||||
|
Extends RTVITextMessageData to include metadata about the output.
|
||||||
|
"""
|
||||||
|
|
||||||
|
spoken: bool = True # Indicates if the text has been spoken by TTS
|
||||||
|
aggregated_by: Optional[Literal["word", "sentence"] | str] = None
|
||||||
|
# Indicates what form the text is in (e.g., by word, sentence, etc.)
|
||||||
|
|
||||||
|
|
||||||
|
class RTVIBotOutputMessage(BaseModel):
|
||||||
|
"""Message containing bot output text.
|
||||||
|
|
||||||
|
An event meant to wholistically represent what the bot is outputting,
|
||||||
|
along with metadata about the output and if it has been spoken.
|
||||||
|
"""
|
||||||
|
|
||||||
|
label: RTVIMessageLiteral = RTVI_MESSAGE_LABEL
|
||||||
|
type: Literal["bot-output"] = "bot-output"
|
||||||
|
data: RTVIBotOutputMessageData
|
||||||
|
|
||||||
|
|
||||||
class RTVIBotTranscriptionMessage(BaseModel):
|
class RTVIBotTranscriptionMessage(BaseModel):
|
||||||
"""Message containing bot transcription text.
|
"""Message containing bot transcription text.
|
||||||
|
|
||||||
@@ -960,6 +983,8 @@ class RTVIObserver(BaseObserver):
|
|||||||
self._last_user_audio_level = 0
|
self._last_user_audio_level = 0
|
||||||
self._last_bot_audio_level = 0
|
self._last_bot_audio_level = 0
|
||||||
|
|
||||||
|
self._skip_tts = None
|
||||||
|
|
||||||
if self._params.system_logs_enabled:
|
if self._params.system_logs_enabled:
|
||||||
self._system_logger_id = logger.add(self._logger_sink)
|
self._system_logger_id = logger.add(self._logger_sink)
|
||||||
|
|
||||||
@@ -1050,8 +1075,7 @@ class RTVIObserver(BaseObserver):
|
|||||||
await self.send_rtvi_message(RTVIBotTTSStoppedMessage())
|
await self.send_rtvi_message(RTVIBotTTSStoppedMessage())
|
||||||
elif isinstance(frame, TTSTextFrame) and self._params.bot_tts_enabled:
|
elif isinstance(frame, TTSTextFrame) and self._params.bot_tts_enabled:
|
||||||
if isinstance(src, BaseOutputTransport):
|
if isinstance(src, BaseOutputTransport):
|
||||||
message = RTVIBotTTSTextMessage(data=RTVITextMessageData(text=frame.text))
|
await self._handle_tts_text_frame(frame)
|
||||||
await self.send_rtvi_message(message)
|
|
||||||
else:
|
else:
|
||||||
mark_as_seen = False
|
mark_as_seen = False
|
||||||
elif isinstance(frame, MetricsFrame) and self._params.metrics_enabled:
|
elif isinstance(frame, MetricsFrame) and self._params.metrics_enabled:
|
||||||
@@ -1115,14 +1139,63 @@ class RTVIObserver(BaseObserver):
|
|||||||
if message:
|
if message:
|
||||||
await self.send_rtvi_message(message)
|
await self.send_rtvi_message(message)
|
||||||
|
|
||||||
|
async def _handle_tts_text_frame(self, frame: TTSTextFrame):
|
||||||
|
"""Handle TTS text output frames."""
|
||||||
|
# send the tts-text message
|
||||||
|
message = RTVIBotTTSTextMessage(data=RTVITextMessageData(text=frame.text))
|
||||||
|
await self.send_rtvi_message(message)
|
||||||
|
# send the bot-output message
|
||||||
|
message = RTVIBotOutputMessage(
|
||||||
|
data=RTVIBotOutputMessageData(
|
||||||
|
text=frame.text, spoken=frame.spoken, aggregated_by=frame.aggregated_by
|
||||||
|
)
|
||||||
|
)
|
||||||
|
await self.send_rtvi_message(message)
|
||||||
|
|
||||||
async def _handle_llm_text_frame(self, frame: LLMTextFrame):
|
async def _handle_llm_text_frame(self, frame: LLMTextFrame):
|
||||||
"""Handle LLM text output frames."""
|
"""Handle LLM text output frames."""
|
||||||
message = RTVIBotLLMTextMessage(data=RTVITextMessageData(text=frame.text))
|
message = RTVIBotLLMTextMessage(data=RTVITextMessageData(text=frame.text))
|
||||||
await self.send_rtvi_message(message)
|
await self.send_rtvi_message(message)
|
||||||
|
|
||||||
|
# initialize skip_tts on first LLMTextFrame
|
||||||
|
if self._skip_tts is None:
|
||||||
|
self._skip_tts = frame.skip_tts
|
||||||
|
|
||||||
|
messages = []
|
||||||
|
should_reset_transcription = False
|
||||||
self._bot_transcription += frame.text
|
self._bot_transcription += frame.text
|
||||||
if match_endofsentence(self._bot_transcription):
|
|
||||||
await self._push_bot_transcription()
|
if not frame.skip_tts and self._skip_tts:
|
||||||
|
# We just switched from skipping TTS to not skipping TTS.
|
||||||
|
# Send and reset any existing transcription.
|
||||||
|
if len(self._bot_transcription) > 0:
|
||||||
|
message.append(
|
||||||
|
RTVIBotOutputMessage(
|
||||||
|
data=RTVIBotOutputMessageData(
|
||||||
|
text=self._bot_transcription, spoken=False, aggregated_by="sentence"
|
||||||
|
)
|
||||||
|
)
|
||||||
|
)
|
||||||
|
should_reset_transcription = True
|
||||||
|
|
||||||
|
if match_endofsentence(self._bot_transcription) and len(self._bot_transcription) > 0:
|
||||||
|
messages.append(
|
||||||
|
RTVIBotTranscriptionMessage(data=RTVITextMessageData(text=self._bot_transcription))
|
||||||
|
)
|
||||||
|
if frame.skip_tts:
|
||||||
|
messages.append(
|
||||||
|
RTVIBotOutputMessage(
|
||||||
|
data=RTVIBotOutputMessageData(
|
||||||
|
text=self._bot_transcription, spoken=False, aggregated_by="sentence"
|
||||||
|
)
|
||||||
|
)
|
||||||
|
)
|
||||||
|
should_reset_transcription = True
|
||||||
|
|
||||||
|
for msg in messages:
|
||||||
|
await self.send_rtvi_message(msg)
|
||||||
|
if should_reset_transcription:
|
||||||
|
self._bot_transcription = ""
|
||||||
|
|
||||||
async def _handle_user_transcriptions(self, frame: Frame):
|
async def _handle_user_transcriptions(self, frame: Frame):
|
||||||
"""Handle user transcription frames."""
|
"""Handle user transcription frames."""
|
||||||
|
|||||||
@@ -1027,7 +1027,7 @@ class AWSNovaSonicLLMService(LLMService):
|
|||||||
logger.debug(f"Assistant response text added: {text}")
|
logger.debug(f"Assistant response text added: {text}")
|
||||||
|
|
||||||
# Report the text of the assistant response.
|
# Report the text of the assistant response.
|
||||||
frame = TTSTextFrame(text)
|
frame = TTSTextFrame(text, aggregated_by="sentence", spoken=True)
|
||||||
frame.includes_inter_frame_spaces = True
|
frame.includes_inter_frame_spaces = True
|
||||||
await self.push_frame(frame)
|
await self.push_frame(frame)
|
||||||
|
|
||||||
@@ -1062,7 +1062,9 @@ class AWSNovaSonicLLMService(LLMService):
|
|||||||
# TTSTextFrame would be ignored otherwise (the interruption frame
|
# TTSTextFrame would be ignored otherwise (the interruption frame
|
||||||
# would have cleared the assistant aggregator state).
|
# would have cleared the assistant aggregator state).
|
||||||
await self.push_frame(LLMFullResponseStartFrame())
|
await self.push_frame(LLMFullResponseStartFrame())
|
||||||
frame = TTSTextFrame(self._assistant_text_buffer)
|
frame = TTSTextFrame(
|
||||||
|
self._assistant_text_buffer, aggregated_by="sentence", spoken=True
|
||||||
|
)
|
||||||
frame.includes_inter_frame_spaces = True
|
frame.includes_inter_frame_spaces = True
|
||||||
await self.push_frame(frame)
|
await self.push_frame(frame)
|
||||||
self._may_need_repush_assistant_text = False
|
self._may_need_repush_assistant_text = False
|
||||||
|
|||||||
@@ -1646,7 +1646,7 @@ class GeminiLiveLLMService(LLMService):
|
|||||||
await self.push_frame(TTSStartedFrame())
|
await self.push_frame(TTSStartedFrame())
|
||||||
await self.push_frame(LLMFullResponseStartFrame())
|
await self.push_frame(LLMFullResponseStartFrame())
|
||||||
|
|
||||||
frame = TTSTextFrame(text=text)
|
frame = TTSTextFrame(text=text, aggregated_by="sentence")
|
||||||
# Gemini Live text already includes any necessary inter-chunk spaces
|
# Gemini Live text already includes any necessary inter-chunk spaces
|
||||||
frame.includes_inter_frame_spaces = True
|
frame.includes_inter_frame_spaces = True
|
||||||
|
|
||||||
|
|||||||
@@ -686,7 +686,7 @@ class OpenAIRealtimeLLMService(LLMService):
|
|||||||
# We receive audio transcript deltas (as opposed to text deltas) when
|
# We receive audio transcript deltas (as opposed to text deltas) when
|
||||||
# the output modality is "audio" (the default)
|
# the output modality is "audio" (the default)
|
||||||
if evt.delta:
|
if evt.delta:
|
||||||
frame = TTSTextFrame(evt.delta)
|
frame = TTSTextFrame(evt.delta, aggregated_by="sentence")
|
||||||
# OpenAI Realtime text already includes any necessary inter-chunk spaces
|
# OpenAI Realtime text already includes any necessary inter-chunk spaces
|
||||||
frame.includes_inter_frame_spaces = True
|
frame.includes_inter_frame_spaces = True
|
||||||
await self.push_frame(frame)
|
await self.push_frame(frame)
|
||||||
|
|||||||
@@ -652,7 +652,7 @@ class OpenAIRealtimeBetaLLMService(LLMService):
|
|||||||
async def _handle_evt_audio_transcript_delta(self, evt):
|
async def _handle_evt_audio_transcript_delta(self, evt):
|
||||||
if evt.delta:
|
if evt.delta:
|
||||||
await self.push_frame(LLMTextFrame(evt.delta))
|
await self.push_frame(LLMTextFrame(evt.delta))
|
||||||
await self.push_frame(TTSTextFrame(evt.delta))
|
await self.push_frame(TTSTextFrame(evt.delta, aggregated_by="sentence", spoken=True))
|
||||||
|
|
||||||
async def _handle_evt_speech_started(self, evt):
|
async def _handle_evt_speech_started(self, evt):
|
||||||
await self._truncate_current_audio_response()
|
await self._truncate_current_audio_response()
|
||||||
|
|||||||
@@ -101,6 +101,8 @@ class TTSService(AIService):
|
|||||||
sample_rate: Optional[int] = None,
|
sample_rate: Optional[int] = None,
|
||||||
# Text aggregator to aggregate incoming tokens and decide when to push to the TTS.
|
# Text aggregator to aggregate incoming tokens and decide when to push to the TTS.
|
||||||
text_aggregator: Optional[BaseTextAggregator] = None,
|
text_aggregator: Optional[BaseTextAggregator] = None,
|
||||||
|
# Types of text aggregations that should not be spoken.
|
||||||
|
skip_aggregator_types: Optional[List[str]] = [],
|
||||||
# Text filter executed after text has been aggregated.
|
# Text filter executed after text has been aggregated.
|
||||||
text_filters: Optional[Sequence[BaseTextFilter]] = None,
|
text_filters: Optional[Sequence[BaseTextFilter]] = None,
|
||||||
text_filter: Optional[BaseTextFilter] = None,
|
text_filter: Optional[BaseTextFilter] = None,
|
||||||
@@ -120,6 +122,7 @@ class TTSService(AIService):
|
|||||||
pause_frame_processing: Whether to pause frame processing during audio generation.
|
pause_frame_processing: Whether to pause frame processing during audio generation.
|
||||||
sample_rate: Output sample rate for generated audio.
|
sample_rate: Output sample rate for generated audio.
|
||||||
text_aggregator: Custom text aggregator for processing incoming text.
|
text_aggregator: Custom text aggregator for processing incoming text.
|
||||||
|
skip_aggregator_types: List of aggregation types that should not be spoken.
|
||||||
text_filters: Sequence of text filters to apply after aggregation.
|
text_filters: Sequence of text filters to apply after aggregation.
|
||||||
text_filter: Single text filter (deprecated, use text_filters).
|
text_filter: Single text filter (deprecated, use text_filters).
|
||||||
|
|
||||||
@@ -142,6 +145,7 @@ class TTSService(AIService):
|
|||||||
self._voice_id: str = ""
|
self._voice_id: str = ""
|
||||||
self._settings: Dict[str, Any] = {}
|
self._settings: Dict[str, Any] = {}
|
||||||
self._text_aggregator: BaseTextAggregator = text_aggregator or SimpleTextAggregator()
|
self._text_aggregator: BaseTextAggregator = text_aggregator or SimpleTextAggregator()
|
||||||
|
self._skip_aggregator_types: List[str] = skip_aggregator_types or []
|
||||||
self._text_filters: Sequence[BaseTextFilter] = text_filters or []
|
self._text_filters: Sequence[BaseTextFilter] = text_filters or []
|
||||||
self._transport_destination: Optional[str] = transport_destination
|
self._transport_destination: Optional[str] = transport_destination
|
||||||
self._tracing_enabled: bool = False
|
self._tracing_enabled: bool = False
|
||||||
@@ -368,10 +372,14 @@ class TTSService(AIService):
|
|||||||
# pause to avoid audio overlapping.
|
# pause to avoid audio overlapping.
|
||||||
await self._maybe_pause_frame_processing()
|
await self._maybe_pause_frame_processing()
|
||||||
|
|
||||||
sentence = self._text_aggregator.text
|
aggregate = self._text_aggregator.text
|
||||||
await self._text_aggregator.reset()
|
await self._text_aggregator.reset()
|
||||||
self._processing_text = False
|
self._processing_text = False
|
||||||
await self._push_tts_frames(sentence)
|
await self._push_tts_frames(
|
||||||
|
text=aggregate.text,
|
||||||
|
should_speak=aggregate.type not in self._skip_aggregator_types,
|
||||||
|
aggregated_by=aggregate.type,
|
||||||
|
)
|
||||||
if isinstance(frame, LLMFullResponseEndFrame):
|
if isinstance(frame, LLMFullResponseEndFrame):
|
||||||
if self._push_text_frames:
|
if self._push_text_frames:
|
||||||
await self.push_frame(frame, direction)
|
await self.push_frame(frame, direction)
|
||||||
@@ -380,7 +388,7 @@ class TTSService(AIService):
|
|||||||
elif isinstance(frame, TTSSpeakFrame):
|
elif isinstance(frame, TTSSpeakFrame):
|
||||||
# Store if we were processing text or not so we can set it back.
|
# Store if we were processing text or not so we can set it back.
|
||||||
processing_text = self._processing_text
|
processing_text = self._processing_text
|
||||||
await self._push_tts_frames(frame.text)
|
await self._push_tts_frames(frame.text, should_speak=True, aggregated_by="word")
|
||||||
# We pause processing incoming frames because we are sending data to
|
# We pause processing incoming frames because we are sending data to
|
||||||
# the TTS. We pause to avoid audio overlapping.
|
# the TTS. We pause to avoid audio overlapping.
|
||||||
await self._maybe_pause_frame_processing()
|
await self._maybe_pause_frame_processing()
|
||||||
@@ -472,42 +480,51 @@ class TTSService(AIService):
|
|||||||
text: Optional[str] = None
|
text: Optional[str] = None
|
||||||
if not self._aggregate_sentences:
|
if not self._aggregate_sentences:
|
||||||
text = frame.text
|
text = frame.text
|
||||||
|
should_speak = True
|
||||||
|
aggregated_by = "token"
|
||||||
else:
|
else:
|
||||||
text = await self._text_aggregator.aggregate(frame.text)
|
aggregate = await self._text_aggregator.aggregate(frame.text)
|
||||||
|
if aggregate:
|
||||||
|
text = aggregate.text
|
||||||
|
should_speak = aggregate.type not in self._skip_aggregator_types
|
||||||
|
aggregated_by = aggregate.type
|
||||||
|
|
||||||
if text:
|
if text:
|
||||||
await self._push_tts_frames(text)
|
logger.trace(f"Pushing TTS frames for text: {text}, {should_speak}, {aggregated_by}")
|
||||||
|
await self._push_tts_frames(text, should_speak, aggregated_by)
|
||||||
|
|
||||||
async def _push_tts_frames(self, text: str):
|
async def _push_tts_frames(self, text: str, should_speak: bool, aggregated_by: str):
|
||||||
# Remove leading newlines only
|
if should_speak:
|
||||||
text = text.lstrip("\n")
|
# Remove leading newlines only
|
||||||
|
text = text.lstrip("\n")
|
||||||
|
|
||||||
# Don't send only whitespace. This causes problems for some TTS models. But also don't
|
# Don't send only whitespace. This causes problems for some TTS models. But also don't
|
||||||
# strip all whitespace, as whitespace can influence prosody.
|
# strip all whitespace, as whitespace can influence prosody.
|
||||||
if not text.strip():
|
if not text.strip():
|
||||||
return
|
return
|
||||||
|
|
||||||
# This is just a flag that indicates if we sent something to the TTS
|
# This is just a flag that indicates if we sent something to the TTS
|
||||||
# service. It will be cleared if we sent text because of a TTSSpeakFrame
|
# service. It will be cleared if we sent text because of a TTSSpeakFrame
|
||||||
# or when we received an LLMFullResponseEndFrame
|
# or when we received an LLMFullResponseEndFrame
|
||||||
self._processing_text = True
|
self._processing_text = True
|
||||||
|
|
||||||
await self.start_processing_metrics()
|
await self.start_processing_metrics()
|
||||||
|
|
||||||
# Process all filter.
|
# Process all filter.
|
||||||
for filter in self._text_filters:
|
for filter in self._text_filters:
|
||||||
await filter.reset_interruption()
|
await filter.reset_interruption()
|
||||||
text = await filter.filter(text)
|
text = await filter.filter(text)
|
||||||
|
|
||||||
if text:
|
if text:
|
||||||
await self.process_generator(self.run_tts(text))
|
await self.push_frame(TTSTextFrame(text, spoken=True, aggregated_by=aggregated_by))
|
||||||
|
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 or not should_speak:
|
||||||
# We send the original text after the audio. This way, if we are
|
# We send the original text after the audio. This way, if we are
|
||||||
# interrupted, the text is not added to the assistant context.
|
# interrupted, the text is not added to the assistant context.
|
||||||
frame = TTSTextFrame(text)
|
frame = TTSTextFrame(text, spoken=should_speak, aggregated_by=aggregated_by)
|
||||||
frame.includes_inter_frame_spaces = self.includes_inter_frame_spaces
|
frame.includes_inter_frame_spaces = self.includes_inter_frame_spaces
|
||||||
await self.push_frame(frame)
|
await self.push_frame(frame)
|
||||||
|
|
||||||
@@ -635,7 +652,7 @@ class WordTTSService(TTSService):
|
|||||||
frame = TTSStoppedFrame()
|
frame = TTSStoppedFrame()
|
||||||
frame.pts = last_pts
|
frame.pts = last_pts
|
||||||
else:
|
else:
|
||||||
frame = TTSTextFrame(word)
|
frame = TTSTextFrame(word, spoken=True, aggregated_by="word")
|
||||||
frame.pts = self._initial_word_timestamp + timestamp
|
frame.pts = self._initial_word_timestamp + timestamp
|
||||||
if frame:
|
if frame:
|
||||||
last_pts = frame.pts
|
last_pts = frame.pts
|
||||||
|
|||||||
@@ -12,9 +12,38 @@ aggregated text should be sent for speech synthesis.
|
|||||||
"""
|
"""
|
||||||
|
|
||||||
from abc import ABC, abstractmethod
|
from abc import ABC, abstractmethod
|
||||||
|
from dataclasses import dataclass
|
||||||
from typing import Optional
|
from typing import Optional
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class Aggregation:
|
||||||
|
"""Data class representing aggregated text and its type.
|
||||||
|
|
||||||
|
An Aggregation object is created whenever a stream of text is aggregated by
|
||||||
|
a text aggregator. It contains the aggregated text and a type indicating
|
||||||
|
the nature of the aggregation.
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(self, text: str, type: str):
|
||||||
|
"""Initialize an aggregation instance.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
text: The aggregated text content.
|
||||||
|
type: The type of aggregation the text represents (e.g., 'sentence', 'word', 'token', 'my_custom_aggregation').
|
||||||
|
"""
|
||||||
|
self.text = text
|
||||||
|
self.type = type
|
||||||
|
|
||||||
|
def __str__(self) -> str:
|
||||||
|
"""Return a string representation of the aggregation.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
A descriptive string showing the type and text of the aggregation.
|
||||||
|
"""
|
||||||
|
return f"Aggregation by {self.type}: {self.text}"
|
||||||
|
|
||||||
|
|
||||||
class BaseTextAggregator(ABC):
|
class BaseTextAggregator(ABC):
|
||||||
"""Base class for text aggregators in the Pipecat framework.
|
"""Base class for text aggregators in the Pipecat framework.
|
||||||
|
|
||||||
@@ -30,7 +59,7 @@ class BaseTextAggregator(ABC):
|
|||||||
|
|
||||||
@property
|
@property
|
||||||
@abstractmethod
|
@abstractmethod
|
||||||
def text(self) -> str:
|
def text(self) -> Aggregation:
|
||||||
"""Get the currently aggregated text.
|
"""Get the currently aggregated text.
|
||||||
|
|
||||||
Subclasses must implement this property to return the text that has
|
Subclasses must implement this property to return the text that has
|
||||||
@@ -42,12 +71,13 @@ class BaseTextAggregator(ABC):
|
|||||||
pass
|
pass
|
||||||
|
|
||||||
@abstractmethod
|
@abstractmethod
|
||||||
async def aggregate(self, text: str) -> Optional[str]:
|
async def aggregate(self, text: str) -> Optional[Aggregation]:
|
||||||
"""Aggregate the specified text with the currently accumulated text.
|
"""Aggregate the specified text with the currently accumulated text.
|
||||||
|
|
||||||
This method should be implemented to define how the new text contributes
|
This method should be implemented to define how the new text contributes
|
||||||
to the aggregation process. It returns the updated aggregated text if
|
to the aggregation process. It returns the aggregated text and a string
|
||||||
it's ready to be processed, or None otherwise.
|
describing how it was aggregated if it's ready to be processed,
|
||||||
|
or None otherwise.
|
||||||
|
|
||||||
Subclasses should implement their specific logic for:
|
Subclasses should implement their specific logic for:
|
||||||
|
|
||||||
|
|||||||
@@ -12,15 +12,15 @@ support for custom handlers and configurable pattern removal.
|
|||||||
"""
|
"""
|
||||||
|
|
||||||
import re
|
import re
|
||||||
from typing import Awaitable, Callable, Optional, Tuple
|
from typing import Awaitable, Callable, List, Optional, Tuple
|
||||||
|
|
||||||
from loguru import logger
|
from loguru import logger
|
||||||
|
|
||||||
from pipecat.utils.string import match_endofsentence
|
from pipecat.utils.string import match_endofsentence
|
||||||
from pipecat.utils.text.base_text_aggregator import BaseTextAggregator
|
from pipecat.utils.text.base_text_aggregator import Aggregation, BaseTextAggregator
|
||||||
|
|
||||||
|
|
||||||
class PatternMatch:
|
class PatternMatch(Aggregation):
|
||||||
"""Represents a matched pattern pair with its content.
|
"""Represents a matched pattern pair with its content.
|
||||||
|
|
||||||
A PatternMatch object is created when a complete pattern pair is found
|
A PatternMatch object is created when a complete pattern pair is found
|
||||||
@@ -29,17 +29,19 @@ class PatternMatch:
|
|||||||
content between the patterns.
|
content between the patterns.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
def __init__(self, pattern_id: str, full_match: str, content: str):
|
def __init__(self, pattern_id: str, full_match: str, content: str, type: str):
|
||||||
"""Initialize a pattern match.
|
"""Initialize a pattern match.
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
pattern_id: The identifier of the matched pattern pair.
|
pattern_id: The identifier of the matched pattern pair.
|
||||||
full_match: The complete text including start and end patterns.
|
full_match: The complete text including start and end patterns.
|
||||||
content: The text content between the start and end patterns.
|
content: The text content between the start and end patterns.
|
||||||
|
type: The type of aggregation the matched content represents
|
||||||
|
(e.g., 'code', 'speaker', 'custom').
|
||||||
"""
|
"""
|
||||||
|
super().__init__(text=content, type=type)
|
||||||
self.pattern_id = pattern_id
|
self.pattern_id = pattern_id
|
||||||
self.full_match = full_match
|
self.full_match = full_match
|
||||||
self.content = content
|
|
||||||
|
|
||||||
def __str__(self) -> str:
|
def __str__(self) -> str:
|
||||||
"""Return a string representation of the pattern match.
|
"""Return a string representation of the pattern match.
|
||||||
@@ -47,7 +49,7 @@ class PatternMatch:
|
|||||||
Returns:
|
Returns:
|
||||||
A descriptive string showing the pattern ID and content.
|
A descriptive string showing the pattern ID and content.
|
||||||
"""
|
"""
|
||||||
return f"PatternMatch(id={self.pattern_id}, content={self.content})"
|
return f"PatternMatch(id={self.pattern_id}, content={self.text}, full_match={self.full_match}, type={self.type})"
|
||||||
|
|
||||||
|
|
||||||
class PatternPairAggregator(BaseTextAggregator):
|
class PatternPairAggregator(BaseTextAggregator):
|
||||||
@@ -64,7 +66,7 @@ class PatternPairAggregator(BaseTextAggregator):
|
|||||||
boundaries.
|
boundaries.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
def __init__(self):
|
def __init__(self, **kwargs):
|
||||||
"""Initialize the pattern pair aggregator.
|
"""Initialize the pattern pair aggregator.
|
||||||
|
|
||||||
Creates an empty aggregator with no patterns or handlers registered.
|
Creates an empty aggregator with no patterns or handlers registered.
|
||||||
@@ -75,16 +77,24 @@ class PatternPairAggregator(BaseTextAggregator):
|
|||||||
self._handlers = {}
|
self._handlers = {}
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def text(self) -> str:
|
def text(self) -> Aggregation:
|
||||||
"""Get the currently buffered text.
|
"""Get the currently aggregated text.
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
The current text buffer content that hasn't been processed yet.
|
The text that has been accumulated in the buffer.
|
||||||
"""
|
"""
|
||||||
return self._text
|
start, curtype = self._match_start_of_pattern(self._text)
|
||||||
|
if curtype:
|
||||||
|
return Aggregation(self._text, curtype)
|
||||||
|
return Aggregation(self._text, "sentence")
|
||||||
|
|
||||||
def add_pattern_pair(
|
def add_pattern_pair(
|
||||||
self, pattern_id: str, start_pattern: str, end_pattern: str, remove_match: bool = True
|
self,
|
||||||
|
pattern_id: str,
|
||||||
|
start_pattern: str,
|
||||||
|
end_pattern: str,
|
||||||
|
type: str,
|
||||||
|
remove_match: bool = True,
|
||||||
) -> "PatternPairAggregator":
|
) -> "PatternPairAggregator":
|
||||||
"""Add a pattern pair to detect in the text.
|
"""Add a pattern pair to detect in the text.
|
||||||
|
|
||||||
@@ -96,7 +106,9 @@ class PatternPairAggregator(BaseTextAggregator):
|
|||||||
pattern_id: Unique identifier for this pattern pair.
|
pattern_id: Unique identifier for this pattern pair.
|
||||||
start_pattern: Pattern that marks the beginning of content.
|
start_pattern: Pattern that marks the beginning of content.
|
||||||
end_pattern: Pattern that marks the end of content.
|
end_pattern: Pattern that marks the end of content.
|
||||||
remove_match: Whether to remove the matched content from the text.
|
type: The type of aggregation the matched content represents
|
||||||
|
(e.g., 'code', 'speaker', 'custom').
|
||||||
|
remove_match: Whether to remove the matched content from the text returned.
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
Self for method chaining.
|
Self for method chaining.
|
||||||
@@ -104,6 +116,7 @@ class PatternPairAggregator(BaseTextAggregator):
|
|||||||
self._patterns[pattern_id] = {
|
self._patterns[pattern_id] = {
|
||||||
"start": start_pattern,
|
"start": start_pattern,
|
||||||
"end": end_pattern,
|
"end": end_pattern,
|
||||||
|
"type": type,
|
||||||
"remove_match": remove_match,
|
"remove_match": remove_match,
|
||||||
}
|
}
|
||||||
return self
|
return self
|
||||||
@@ -127,7 +140,7 @@ class PatternPairAggregator(BaseTextAggregator):
|
|||||||
self._handlers[pattern_id] = handler
|
self._handlers[pattern_id] = handler
|
||||||
return self
|
return self
|
||||||
|
|
||||||
async def _process_complete_patterns(self, text: str) -> Tuple[str, bool]:
|
async def _process_complete_patterns(self, text: str) -> Tuple[List[PatternMatch], str]:
|
||||||
"""Process all complete pattern pairs in the text.
|
"""Process all complete pattern pairs in the text.
|
||||||
|
|
||||||
Searches for all complete pattern pairs in the text, calls the
|
Searches for all complete pattern pairs in the text, calls the
|
||||||
@@ -137,19 +150,20 @@ class PatternPairAggregator(BaseTextAggregator):
|
|||||||
text: The text to process for pattern matches.
|
text: The text to process for pattern matches.
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
Tuple of (processed_text, was_modified) where:
|
Tuple of (all_matches, processed_text) where:
|
||||||
|
|
||||||
- processed_text is the text after processing patterns
|
- all_matches is a list of all pattern matches found. Note: There really should only ever be 1.
|
||||||
- was_modified indicates whether any changes were made
|
- processed_text is the text after processing patterns. If no patterns are found, it will be the same as input text.
|
||||||
"""
|
"""
|
||||||
|
all_matches = []
|
||||||
processed_text = text
|
processed_text = text
|
||||||
modified = False
|
|
||||||
|
|
||||||
for pattern_id, pattern_info in self._patterns.items():
|
for pattern_id, pattern_info in self._patterns.items():
|
||||||
# Escape special regex characters in the patterns
|
# Escape special regex characters in the patterns
|
||||||
start = re.escape(pattern_info["start"])
|
start = re.escape(pattern_info["start"])
|
||||||
end = re.escape(pattern_info["end"])
|
end = re.escape(pattern_info["end"])
|
||||||
remove_match = pattern_info["remove_match"]
|
remove_match = pattern_info["remove_match"]
|
||||||
|
match_type = pattern_info["type"]
|
||||||
|
|
||||||
# Create regex to match from start pattern to end pattern
|
# Create regex to match from start pattern to end pattern
|
||||||
# The .*? is non-greedy to handle nested patterns
|
# The .*? is non-greedy to handle nested patterns
|
||||||
@@ -165,7 +179,7 @@ class PatternPairAggregator(BaseTextAggregator):
|
|||||||
|
|
||||||
# Create pattern match object
|
# Create pattern match object
|
||||||
pattern_match = PatternMatch(
|
pattern_match = PatternMatch(
|
||||||
pattern_id=pattern_id, full_match=full_match, content=content
|
pattern_id=pattern_id, full_match=full_match, content=content, type=match_type
|
||||||
)
|
)
|
||||||
|
|
||||||
# Call the appropriate handler if registered
|
# Call the appropriate handler if registered
|
||||||
@@ -178,11 +192,13 @@ class PatternPairAggregator(BaseTextAggregator):
|
|||||||
# Remove the pattern from the text if configured
|
# Remove the pattern from the text if configured
|
||||||
if remove_match:
|
if remove_match:
|
||||||
processed_text = processed_text.replace(full_match, "", 1)
|
processed_text = processed_text.replace(full_match, "", 1)
|
||||||
modified = True
|
# modified = True
|
||||||
|
else:
|
||||||
|
all_matches.append(pattern_match)
|
||||||
|
|
||||||
return processed_text, modified
|
return all_matches, processed_text
|
||||||
|
|
||||||
def _has_incomplete_patterns(self, text: str) -> bool:
|
def _match_start_of_pattern(self, text: str) -> Optional[Tuple[int, str]]:
|
||||||
"""Check if text contains incomplete pattern pairs.
|
"""Check if text contains incomplete pattern pairs.
|
||||||
|
|
||||||
Determines whether the text contains any start patterns without
|
Determines whether the text contains any start patterns without
|
||||||
@@ -192,7 +208,8 @@ class PatternPairAggregator(BaseTextAggregator):
|
|||||||
text: The text to check for incomplete patterns.
|
text: The text to check for incomplete patterns.
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
True if there are incomplete patterns, False otherwise.
|
A tuple of (start_index, type) if an incomplete pattern is found,
|
||||||
|
or None if no patterns are found or all patterns are complete.
|
||||||
"""
|
"""
|
||||||
for pattern_id, pattern_info in self._patterns.items():
|
for pattern_id, pattern_info in self._patterns.items():
|
||||||
start = pattern_info["start"]
|
start = pattern_info["start"]
|
||||||
@@ -203,12 +220,16 @@ class PatternPairAggregator(BaseTextAggregator):
|
|||||||
end_count = text.count(end)
|
end_count = text.count(end)
|
||||||
|
|
||||||
# If there are more starts than ends, we have incomplete patterns
|
# If there are more starts than ends, we have incomplete patterns
|
||||||
|
# Again, this is written generically but there only ever should
|
||||||
|
# be one pattern active at a time, so the counts should be 0 or 1.
|
||||||
|
# Which is why we base the return on the first found.
|
||||||
if start_count > end_count:
|
if start_count > end_count:
|
||||||
return True
|
start_index = text.find(start)
|
||||||
|
return [start_index, pattern_info["type"]]
|
||||||
|
|
||||||
return False
|
return None, None
|
||||||
|
|
||||||
async def aggregate(self, text: str) -> Optional[str]:
|
async def aggregate(self, text: str) -> Optional[PatternMatch]:
|
||||||
"""Aggregate text and process pattern pairs.
|
"""Aggregate text and process pattern pairs.
|
||||||
|
|
||||||
This method adds the new text to the buffer, processes any complete pattern
|
This method adds the new text to the buffer, processes any complete pattern
|
||||||
@@ -227,16 +248,28 @@ class PatternPairAggregator(BaseTextAggregator):
|
|||||||
self._text += text
|
self._text += text
|
||||||
|
|
||||||
# Process any complete patterns in the buffer
|
# Process any complete patterns in the buffer
|
||||||
processed_text, modified = await self._process_complete_patterns(self._text)
|
patterns, processed_text = await self._process_complete_patterns(self._text)
|
||||||
|
|
||||||
# Only update the buffer if modifications were made
|
self._text = processed_text
|
||||||
if modified:
|
|
||||||
self._text = processed_text
|
#
|
||||||
|
if len(patterns) > 0:
|
||||||
|
if len(patterns) > 1:
|
||||||
|
logger.warning(
|
||||||
|
f"Multiple patterns matched: {[p.pattern_id for p in patterns]}. Only the first pattern will be returned."
|
||||||
|
)
|
||||||
|
self._text = ""
|
||||||
|
return patterns[0]
|
||||||
|
|
||||||
# Check if we have incomplete patterns
|
# Check if we have incomplete patterns
|
||||||
if self._has_incomplete_patterns(self._text):
|
start, curtype = self._match_start_of_pattern(self._text)
|
||||||
|
if start is not None:
|
||||||
# Still waiting for complete patterns
|
# Still waiting for complete patterns
|
||||||
return None
|
if start == 0:
|
||||||
|
return None
|
||||||
|
result = self._text[:start]
|
||||||
|
self._text = self._text[start:]
|
||||||
|
return PatternMatch(f"_sentence", result, result, "sentence")
|
||||||
|
|
||||||
# Find sentence boundary if no incomplete patterns
|
# Find sentence boundary if no incomplete patterns
|
||||||
eos_marker = match_endofsentence(self._text)
|
eos_marker = match_endofsentence(self._text)
|
||||||
@@ -244,7 +277,7 @@ class PatternPairAggregator(BaseTextAggregator):
|
|||||||
# Extract text up to the sentence boundary
|
# Extract text up to the sentence boundary
|
||||||
result = self._text[:eos_marker]
|
result = self._text[:eos_marker]
|
||||||
self._text = self._text[eos_marker:]
|
self._text = self._text[eos_marker:]
|
||||||
return result
|
return PatternMatch(f"_sentence", result, result, "sentence")
|
||||||
|
|
||||||
# No complete sentence found yet
|
# No complete sentence found yet
|
||||||
return None
|
return None
|
||||||
|
|||||||
@@ -14,7 +14,7 @@ text processing scenarios.
|
|||||||
from typing import Optional
|
from typing import Optional
|
||||||
|
|
||||||
from pipecat.utils.string import match_endofsentence
|
from pipecat.utils.string import match_endofsentence
|
||||||
from pipecat.utils.text.base_text_aggregator import BaseTextAggregator
|
from pipecat.utils.text.base_text_aggregator import Aggregation, BaseTextAggregator
|
||||||
|
|
||||||
|
|
||||||
class SimpleTextAggregator(BaseTextAggregator):
|
class SimpleTextAggregator(BaseTextAggregator):
|
||||||
@@ -33,15 +33,15 @@ class SimpleTextAggregator(BaseTextAggregator):
|
|||||||
self._text = ""
|
self._text = ""
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def text(self) -> str:
|
def text(self) -> Aggregation:
|
||||||
"""Get the currently aggregated text.
|
"""Get the currently aggregated text.
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
The text that has been accumulated in the buffer.
|
The text that has been accumulated in the buffer.
|
||||||
"""
|
"""
|
||||||
return self._text
|
return Aggregation(self._text, "sentence")
|
||||||
|
|
||||||
async def aggregate(self, text: str) -> Optional[str]:
|
async def aggregate(self, text: str) -> Optional[Aggregation]:
|
||||||
"""Aggregate text and return completed sentences.
|
"""Aggregate text and return completed sentences.
|
||||||
|
|
||||||
Adds the new text to the buffer and checks for end-of-sentence markers.
|
Adds the new text to the buffer and checks for end-of-sentence markers.
|
||||||
@@ -64,7 +64,7 @@ class SimpleTextAggregator(BaseTextAggregator):
|
|||||||
result = self._text[:eos_end_marker]
|
result = self._text[:eos_end_marker]
|
||||||
self._text = self._text[eos_end_marker:]
|
self._text = self._text[eos_end_marker:]
|
||||||
|
|
||||||
return result
|
return Aggregation(result, "sentence") if result else None
|
||||||
|
|
||||||
async def handle_interruption(self):
|
async def handle_interruption(self):
|
||||||
"""Handle interruptions by clearing the text buffer.
|
"""Handle interruptions by clearing the text buffer.
|
||||||
|
|||||||
@@ -14,7 +14,7 @@ as a unit regardless of internal punctuation.
|
|||||||
from typing import Optional, Sequence
|
from typing import Optional, Sequence
|
||||||
|
|
||||||
from pipecat.utils.string import StartEndTags, match_endofsentence, parse_start_end_tags
|
from pipecat.utils.string import StartEndTags, match_endofsentence, parse_start_end_tags
|
||||||
from pipecat.utils.text.base_text_aggregator import BaseTextAggregator
|
from pipecat.utils.text.base_text_aggregator import Aggregation, BaseTextAggregator
|
||||||
|
|
||||||
|
|
||||||
class SkipTagsAggregator(BaseTextAggregator):
|
class SkipTagsAggregator(BaseTextAggregator):
|
||||||
@@ -49,9 +49,9 @@ class SkipTagsAggregator(BaseTextAggregator):
|
|||||||
Returns:
|
Returns:
|
||||||
The current text buffer content that hasn't been processed yet.
|
The current text buffer content that hasn't been processed yet.
|
||||||
"""
|
"""
|
||||||
return self._text
|
return Aggregation(self._text, "sentence")
|
||||||
|
|
||||||
async def aggregate(self, text: str) -> Optional[str]:
|
async def aggregate(self, text: str) -> Optional[Aggregation]:
|
||||||
"""Aggregate text while respecting tag boundaries.
|
"""Aggregate text while respecting tag boundaries.
|
||||||
|
|
||||||
This method adds the new text to the buffer, processes any complete
|
This method adds the new text to the buffer, processes any complete
|
||||||
@@ -80,7 +80,7 @@ class SkipTagsAggregator(BaseTextAggregator):
|
|||||||
# Extract text up to the sentence boundary
|
# Extract text up to the sentence boundary
|
||||||
result = self._text[:eos_marker]
|
result = self._text[:eos_marker]
|
||||||
self._text = self._text[eos_marker:]
|
self._text = self._text[eos_marker:]
|
||||||
return result
|
return Aggregation(result, "sentence")
|
||||||
|
|
||||||
# No complete sentence found yet
|
# No complete sentence found yet
|
||||||
return None
|
return None
|
||||||
|
|||||||
@@ -130,11 +130,11 @@ class TestUserTranscriptProcessor(unittest.IsolatedAsyncioTestCase):
|
|||||||
frames_to_send = [
|
frames_to_send = [
|
||||||
BotStartedSpeakingFrame(),
|
BotStartedSpeakingFrame(),
|
||||||
SleepFrame(), # Wait for StartedSpeaking to process
|
SleepFrame(), # Wait for StartedSpeaking to process
|
||||||
TTSTextFrame(text="Hello"),
|
TTSTextFrame(text="Hello", aggregated_by="word"),
|
||||||
TTSTextFrame(text="world!"),
|
TTSTextFrame(text="world!", aggregated_by="word"),
|
||||||
TTSTextFrame(text="How"),
|
TTSTextFrame(text="How", aggregated_by="word"),
|
||||||
TTSTextFrame(text="are"),
|
TTSTextFrame(text="are", aggregated_by="word"),
|
||||||
TTSTextFrame(text="you?"),
|
TTSTextFrame(text="you?", aggregated_by="word"),
|
||||||
SleepFrame(), # Wait for text frames to queue
|
SleepFrame(), # Wait for text frames to queue
|
||||||
BotStoppedSpeakingFrame(),
|
BotStoppedSpeakingFrame(),
|
||||||
]
|
]
|
||||||
@@ -195,9 +195,9 @@ class TestUserTranscriptProcessor(unittest.IsolatedAsyncioTestCase):
|
|||||||
frames_to_send = [
|
frames_to_send = [
|
||||||
BotStartedSpeakingFrame(),
|
BotStartedSpeakingFrame(),
|
||||||
SleepFrame(),
|
SleepFrame(),
|
||||||
TTSTextFrame(text=""), # Empty text
|
TTSTextFrame(text="", aggregated_by="word"), # Empty text
|
||||||
TTSTextFrame(text=" "), # Just whitespace
|
TTSTextFrame(text=" ", aggregated_by="word"), # Just whitespace
|
||||||
TTSTextFrame(text="\n"), # Just newline
|
TTSTextFrame(text="\n", aggregated_by="word"), # Just newline
|
||||||
BotStoppedSpeakingFrame(),
|
BotStoppedSpeakingFrame(),
|
||||||
# Pipeline ends here; run_test will automatically send EndFrame
|
# Pipeline ends here; run_test will automatically send EndFrame
|
||||||
]
|
]
|
||||||
@@ -235,14 +235,14 @@ class TestUserTranscriptProcessor(unittest.IsolatedAsyncioTestCase):
|
|||||||
frames_to_send = [
|
frames_to_send = [
|
||||||
BotStartedSpeakingFrame(),
|
BotStartedSpeakingFrame(),
|
||||||
SleepFrame(),
|
SleepFrame(),
|
||||||
TTSTextFrame(text="Hello"),
|
TTSTextFrame(text="Hello", aggregated_by="word"),
|
||||||
TTSTextFrame(text="world!"),
|
TTSTextFrame(text="world!", aggregated_by="word"),
|
||||||
SleepFrame(),
|
SleepFrame(),
|
||||||
InterruptionFrame(), # User interrupts here
|
InterruptionFrame(), # User interrupts here
|
||||||
SleepFrame(),
|
SleepFrame(),
|
||||||
BotStartedSpeakingFrame(),
|
BotStartedSpeakingFrame(),
|
||||||
TTSTextFrame(text="New"),
|
TTSTextFrame(text="New", aggregated_by="word"),
|
||||||
TTSTextFrame(text="response"),
|
TTSTextFrame(text="response", aggregated_by="word"),
|
||||||
SleepFrame(),
|
SleepFrame(),
|
||||||
BotStoppedSpeakingFrame(),
|
BotStoppedSpeakingFrame(),
|
||||||
]
|
]
|
||||||
@@ -299,8 +299,8 @@ class TestUserTranscriptProcessor(unittest.IsolatedAsyncioTestCase):
|
|||||||
frames_to_send = [
|
frames_to_send = [
|
||||||
BotStartedSpeakingFrame(),
|
BotStartedSpeakingFrame(),
|
||||||
SleepFrame(),
|
SleepFrame(),
|
||||||
TTSTextFrame(text="Hello"),
|
TTSTextFrame(text="Hello", aggregated_by="word"),
|
||||||
TTSTextFrame(text="world"),
|
TTSTextFrame(text="world", aggregated_by="word"),
|
||||||
# Pipeline ends here; run_test will automatically send EndFrame
|
# Pipeline ends here; run_test will automatically send EndFrame
|
||||||
]
|
]
|
||||||
|
|
||||||
@@ -338,8 +338,8 @@ class TestUserTranscriptProcessor(unittest.IsolatedAsyncioTestCase):
|
|||||||
frames_to_send = [
|
frames_to_send = [
|
||||||
BotStartedSpeakingFrame(),
|
BotStartedSpeakingFrame(),
|
||||||
SleepFrame(),
|
SleepFrame(),
|
||||||
TTSTextFrame(text="Hello"),
|
TTSTextFrame(text="Hello", aggregated_by="word"),
|
||||||
TTSTextFrame(text="world"),
|
TTSTextFrame(text="world", aggregated_by="word"),
|
||||||
SleepFrame(), # Ensure messages are processed
|
SleepFrame(), # Ensure messages are processed
|
||||||
CancelFrame(),
|
CancelFrame(),
|
||||||
]
|
]
|
||||||
@@ -401,8 +401,8 @@ class TestUserTranscriptProcessor(unittest.IsolatedAsyncioTestCase):
|
|||||||
frames_to_send = [
|
frames_to_send = [
|
||||||
BotStartedSpeakingFrame(),
|
BotStartedSpeakingFrame(),
|
||||||
SleepFrame(),
|
SleepFrame(),
|
||||||
TTSTextFrame(text="Assistant"),
|
TTSTextFrame(text="Assistant", aggregated_by="word"),
|
||||||
TTSTextFrame(text="message"),
|
TTSTextFrame(text="message", aggregated_by="word"),
|
||||||
BotStoppedSpeakingFrame(),
|
BotStoppedSpeakingFrame(),
|
||||||
]
|
]
|
||||||
|
|
||||||
@@ -439,7 +439,7 @@ class TestUserTranscriptProcessor(unittest.IsolatedAsyncioTestCase):
|
|||||||
|
|
||||||
# Test the specific pattern shared
|
# Test the specific pattern shared
|
||||||
def make_tts_text_frame(text: str) -> TTSTextFrame:
|
def make_tts_text_frame(text: str) -> TTSTextFrame:
|
||||||
frame = TTSTextFrame(text=text)
|
frame = TTSTextFrame(text=text, aggregated_by="word")
|
||||||
frame.includes_inter_frame_spaces = True
|
frame.includes_inter_frame_spaces = True
|
||||||
return frame
|
return frame
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user