Augmented PatternPairAggregator so that matched patterns can...
be treated as their own aggregation, taking advantage of the new ability to assign a type to an aggregation
This commit is contained in:
committed by
Mattie Ruth
parent
dcc20f86e1
commit
24266c238f
27
CHANGELOG.md
27
CHANGELOG.md
@@ -108,6 +108,33 @@ Croatian, Hungarian, Malay, Norwegian, Nynorsk, Slovak, Slovenian, Swedish, and
|
|||||||
produce/consume `Aggregation` objects.
|
produce/consume `Aggregation` objects.
|
||||||
- All uses of the above Aggregators have been updated accordingly.
|
- All uses of the above Aggregators have been updated accordingly.
|
||||||
|
|
||||||
|
- Augmented the `PatternPairAggregator` so that matched patterns can be treated as their own
|
||||||
|
aggregation, taking advantage of the new. To that end:
|
||||||
|
- Introduced a new, preferred version of `add_pattern` to support a new option for treating a
|
||||||
|
match as a separate aggregation returned from `aggregate()`. This replaces the now
|
||||||
|
deprecated `add_pattern_pair` method and you provide a `MatchAction` in lieu of the `remove_match` field.
|
||||||
|
- `MatchAction` enum: `REMOVE`, `KEEP`, `AGGREGATE`, allowing customization for how
|
||||||
|
a match should be handled.
|
||||||
|
- `REMOVE`: The text along with its delimiters will be removed from the streaming text.
|
||||||
|
Sentence aggregation will continue on as if this text did not exist.
|
||||||
|
- `KEEP`: The delimiters will be removed, but the content between them will be kept.
|
||||||
|
Sentence aggregation will continue on with the internal text included.
|
||||||
|
- `AGGREGATE`: The delimiters will be removed and the content between will be treated
|
||||||
|
as a separate aggregation. Any text before the start of the pattern will be
|
||||||
|
returned early, whether or not a complete sentence was found. Then the pattern
|
||||||
|
will be returned. Then the aggregation will continue on sentence matching after
|
||||||
|
the closing delimiter is found. The content between the delimiters is not
|
||||||
|
aggregated by sentence. It is aggregated as one single block of text.
|
||||||
|
- `PatternMatch` now extends `Aggregation` and provides richer info to handlers.
|
||||||
|
- **BREAKING**: The `PatternMatch` type returned to handlers registered via `on_pattern_match`
|
||||||
|
has been updated to subclass from the new `Aggregation` type, which means that `content`
|
||||||
|
has been replaced with `text` and `pattern_id` has been replaced with `type`:
|
||||||
|
```
|
||||||
|
async dev on_match_tag(match: PatternMatch):
|
||||||
|
pattern = match.type # instead of match.pattern_id
|
||||||
|
text = match.text # instead of match.content
|
||||||
|
```
|
||||||
|
|
||||||
### Deprecated
|
### Deprecated
|
||||||
|
|
||||||
- The `api_key` parameter in `GeminiTTSService` is deprecated. Use
|
- The `api_key` parameter in `GeminiTTSService` is deprecated. Use
|
||||||
|
|||||||
@@ -62,7 +62,11 @@ from pipecat.services.openai.llm import OpenAILLMService
|
|||||||
from pipecat.transports.base_transport import BaseTransport, TransportParams
|
from pipecat.transports.base_transport import BaseTransport, TransportParams
|
||||||
from pipecat.transports.daily.transport import DailyParams
|
from pipecat.transports.daily.transport import DailyParams
|
||||||
from pipecat.transports.websocket.fastapi import FastAPIWebsocketParams
|
from pipecat.transports.websocket.fastapi import FastAPIWebsocketParams
|
||||||
from pipecat.utils.text.pattern_pair_aggregator import PatternMatch, PatternPairAggregator
|
from pipecat.utils.text.pattern_pair_aggregator import (
|
||||||
|
MatchAction,
|
||||||
|
PatternMatch,
|
||||||
|
PatternPairAggregator,
|
||||||
|
)
|
||||||
|
|
||||||
load_dotenv(override=True)
|
load_dotenv(override=True)
|
||||||
|
|
||||||
@@ -106,16 +110,16 @@ async def run_bot(transport: BaseTransport, runner_args: RunnerArguments):
|
|||||||
pattern_aggregator = PatternPairAggregator()
|
pattern_aggregator = PatternPairAggregator()
|
||||||
|
|
||||||
# Add pattern for voice switching
|
# Add pattern for voice switching
|
||||||
pattern_aggregator.add_pattern_pair(
|
pattern_aggregator.add_pattern(
|
||||||
pattern_id="voice_tag",
|
type="voice",
|
||||||
start_pattern="<voice>",
|
start_pattern="<voice>",
|
||||||
end_pattern="</voice>",
|
end_pattern="</voice>",
|
||||||
remove_match=True,
|
action=MatchAction.REMOVE, # Remove tags from final text
|
||||||
)
|
)
|
||||||
|
|
||||||
# Register handler for voice switching
|
# Register handler for voice switching
|
||||||
async def on_voice_tag(match: PatternMatch):
|
async def on_voice_tag(match: PatternMatch):
|
||||||
voice_name = match.content.strip().lower()
|
voice_name = match.text.strip().lower()
|
||||||
if voice_name in VOICE_IDS:
|
if voice_name in VOICE_IDS:
|
||||||
# First flush any existing audio to finish the current context
|
# First flush any existing audio to finish the current context
|
||||||
await tts.flush_audio()
|
await tts.flush_audio()
|
||||||
@@ -125,7 +129,7 @@ async def run_bot(transport: BaseTransport, runner_args: RunnerArguments):
|
|||||||
else:
|
else:
|
||||||
logger.warning(f"Unknown voice: {voice_name}")
|
logger.warning(f"Unknown voice: {voice_name}")
|
||||||
|
|
||||||
pattern_aggregator.on_pattern_match("voice_tag", on_voice_tag)
|
pattern_aggregator.on_pattern_match("voice", on_voice_tag)
|
||||||
|
|
||||||
stt = DeepgramSTTService(api_key=os.getenv("DEEPGRAM_API_KEY"))
|
stt = DeepgramSTTService(api_key=os.getenv("DEEPGRAM_API_KEY"))
|
||||||
|
|
||||||
|
|||||||
@@ -31,7 +31,11 @@ from pipecat.pipeline.pipeline import Pipeline
|
|||||||
from pipecat.processors.aggregators.openai_llm_context import OpenAILLMContextFrame
|
from pipecat.processors.aggregators.openai_llm_context import OpenAILLMContextFrame
|
||||||
from pipecat.processors.frame_processor import FrameDirection, FrameProcessor
|
from pipecat.processors.frame_processor import FrameDirection, FrameProcessor
|
||||||
from pipecat.services.llm_service import LLMService
|
from pipecat.services.llm_service import LLMService
|
||||||
from pipecat.utils.text.pattern_pair_aggregator import PatternMatch, PatternPairAggregator
|
from pipecat.utils.text.pattern_pair_aggregator import (
|
||||||
|
MatchAction,
|
||||||
|
PatternMatch,
|
||||||
|
PatternPairAggregator,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
class IVRStatus(Enum):
|
class IVRStatus(Enum):
|
||||||
@@ -114,15 +118,15 @@ class IVRProcessor(FrameProcessor):
|
|||||||
def _setup_xml_patterns(self):
|
def _setup_xml_patterns(self):
|
||||||
"""Set up XML pattern detection and handlers."""
|
"""Set up XML pattern detection and handlers."""
|
||||||
# Register DTMF pattern
|
# Register DTMF pattern
|
||||||
self._aggregator.add_pattern_pair("dtmf", "<dtmf>", "</dtmf>", remove_match=True)
|
self._aggregator.add_pattern("dtmf", "<dtmf>", "</dtmf>", action=MatchAction.REMOVE)
|
||||||
self._aggregator.on_pattern_match("dtmf", self._handle_dtmf_action)
|
self._aggregator.on_pattern_match("dtmf", self._handle_dtmf_action)
|
||||||
|
|
||||||
# Register mode pattern
|
# Register mode pattern
|
||||||
self._aggregator.add_pattern_pair("mode", "<mode>", "</mode>", remove_match=True)
|
self._aggregator.add_pattern("mode", "<mode>", "</mode>", action=MatchAction.REMOVE)
|
||||||
self._aggregator.on_pattern_match("mode", self._handle_mode_action)
|
self._aggregator.on_pattern_match("mode", self._handle_mode_action)
|
||||||
|
|
||||||
# Register IVR pattern
|
# Register IVR pattern
|
||||||
self._aggregator.add_pattern_pair("ivr", "<ivr>", "</ivr>", remove_match=True)
|
self._aggregator.add_pattern("ivr", "<ivr>", "</ivr>", action=MatchAction.REMOVE)
|
||||||
self._aggregator.on_pattern_match("ivr", self._handle_ivr_action)
|
self._aggregator.on_pattern_match("ivr", self._handle_ivr_action)
|
||||||
|
|
||||||
async def process_frame(self, frame: Frame, direction: FrameDirection):
|
async def process_frame(self, frame: Frame, direction: FrameDirection):
|
||||||
@@ -159,7 +163,7 @@ class IVRProcessor(FrameProcessor):
|
|||||||
Args:
|
Args:
|
||||||
match: The pattern match containing DTMF content.
|
match: The pattern match containing DTMF content.
|
||||||
"""
|
"""
|
||||||
value = match.content
|
value = match.text
|
||||||
logger.debug(f"DTMF detected: {value}")
|
logger.debug(f"DTMF detected: {value}")
|
||||||
|
|
||||||
try:
|
try:
|
||||||
@@ -180,7 +184,7 @@ class IVRProcessor(FrameProcessor):
|
|||||||
Args:
|
Args:
|
||||||
match: The pattern match containing IVR status content.
|
match: The pattern match containing IVR status content.
|
||||||
"""
|
"""
|
||||||
status = match.content
|
status = match.text
|
||||||
logger.trace(f"IVR status detected: {status}")
|
logger.trace(f"IVR status detected: {status}")
|
||||||
|
|
||||||
# Convert string to enum, with validation
|
# Convert string to enum, with validation
|
||||||
@@ -211,7 +215,7 @@ class IVRProcessor(FrameProcessor):
|
|||||||
Args:
|
Args:
|
||||||
match: The pattern match containing mode content.
|
match: The pattern match containing mode content.
|
||||||
"""
|
"""
|
||||||
mode = match.content
|
mode = match.text
|
||||||
logger.debug(f"Mode detected: {mode}")
|
logger.debug(f"Mode detected: {mode}")
|
||||||
if mode == "conversation":
|
if mode == "conversation":
|
||||||
await self._handle_conversation()
|
await self._handle_conversation()
|
||||||
|
|||||||
@@ -8,11 +8,12 @@
|
|||||||
|
|
||||||
This module provides an aggregator that identifies and processes content between
|
This module provides an aggregator that identifies and processes content between
|
||||||
pattern pairs (like XML tags or custom delimiters) in streaming text, with
|
pattern pairs (like XML tags or custom delimiters) in streaming text, with
|
||||||
support for custom handlers and configurable pattern removal.
|
support for custom handlers and configurable actions for when a pattern is found.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
import re
|
import re
|
||||||
from typing import Awaitable, Callable, Optional, Tuple
|
from enum import Enum
|
||||||
|
from typing import Awaitable, Callable, List, Optional, Tuple
|
||||||
|
|
||||||
from loguru import logger
|
from loguru import logger
|
||||||
|
|
||||||
@@ -20,7 +21,28 @@ from pipecat.utils.string import match_endofsentence
|
|||||||
from pipecat.utils.text.base_text_aggregator import Aggregation, AggregationType, BaseTextAggregator
|
from pipecat.utils.text.base_text_aggregator import Aggregation, AggregationType, BaseTextAggregator
|
||||||
|
|
||||||
|
|
||||||
class PatternMatch:
|
class MatchAction(Enum):
|
||||||
|
"""Actions to take when a pattern pair is matched.
|
||||||
|
|
||||||
|
Parameters:
|
||||||
|
REMOVE: The text along with its delimiters will be removed from the streaming text.
|
||||||
|
Sentence aggregation will continue on as if this text did not exist.
|
||||||
|
KEEP: The delimiters will be removed, but the content between them will be kept.
|
||||||
|
Sentence aggregation will continue on with the internal text included.
|
||||||
|
AGGREGATE: The delimiters will be removed and the content between will be treated
|
||||||
|
as a separate aggregation. Any text before the start of the pattern will be
|
||||||
|
returned early, whether or not a complete sentence was found. Then the pattern
|
||||||
|
will be returned. Then the aggregation will continue on sentence matching after
|
||||||
|
the closing delimiter is found. The content between the delimiters is not
|
||||||
|
aggregated by sentence. It is aggregated as one single block of text.
|
||||||
|
"""
|
||||||
|
|
||||||
|
REMOVE = "remove"
|
||||||
|
KEEP = "keep"
|
||||||
|
AGGREGATE = "aggregate"
|
||||||
|
|
||||||
|
|
||||||
|
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,25 +51,25 @@ class PatternMatch:
|
|||||||
content between the patterns.
|
content between the patterns.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
def __init__(self, pattern_id: str, full_match: str, content: str):
|
def __init__(self, content: str, type: str, full_match: str):
|
||||||
"""Initialize a pattern match.
|
"""Initialize a pattern match.
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
pattern_id: The identifier of the matched pattern pair.
|
type: The type of the matched pattern pair. It should be representative
|
||||||
|
of the content type (e.g., 'sentence', 'code', 'speaker', 'custom').
|
||||||
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.
|
||||||
"""
|
"""
|
||||||
self.pattern_id = pattern_id
|
super().__init__(text=content, type=type)
|
||||||
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.
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
A descriptive string showing the pattern ID and content.
|
A descriptive string showing the pattern type and content.
|
||||||
"""
|
"""
|
||||||
return f"PatternMatch(id={self.pattern_id}, content={self.content})"
|
return f"PatternMatch(type={self.type}, text={self.text}, full_match={self.full_match})"
|
||||||
|
|
||||||
|
|
||||||
class PatternPairAggregator(BaseTextAggregator):
|
class PatternPairAggregator(BaseTextAggregator):
|
||||||
@@ -55,16 +77,21 @@ class PatternPairAggregator(BaseTextAggregator):
|
|||||||
|
|
||||||
This aggregator buffers text until it can identify complete pattern pairs
|
This aggregator buffers text until it can identify complete pattern pairs
|
||||||
(defined by start and end patterns), processes the content between these
|
(defined by start and end patterns), processes the content between these
|
||||||
patterns using registered handlers, and returns text at sentence boundaries.
|
patterns using registered handlers. By default, its aggregation method
|
||||||
It's particularly useful for processing structured content in streaming text,
|
returns text at sentence boundaries, and remove the content found between
|
||||||
such as XML tags, markdown formatting, or custom delimiters.
|
any matched patterns. However, matched patterns can also be configured to
|
||||||
|
returned as a separate aggregation object containing the content between
|
||||||
|
their start and end patterns or left in, so that only the delimiters are
|
||||||
|
removed and a callback can be triggered.
|
||||||
|
|
||||||
|
This aggregator is particularly useful for processing structured content in
|
||||||
|
streaming text, such as XML tags, markdown formatting, or custom delimiters.
|
||||||
|
|
||||||
The aggregator ensures that patterns spanning multiple text chunks are
|
The aggregator ensures that patterns spanning multiple text chunks are
|
||||||
correctly identified and handles cases where patterns contain sentence
|
correctly identified.
|
||||||
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.
|
||||||
@@ -76,15 +103,26 @@ class PatternPairAggregator(BaseTextAggregator):
|
|||||||
|
|
||||||
@property
|
@property
|
||||||
def text(self) -> Aggregation:
|
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 Aggregation(text=self._text.strip(), type=AggregationType.SENTENCE)
|
pattern_start = self._match_start_of_pattern(self._text)
|
||||||
|
stripped_text = self._text.strip()
|
||||||
|
type = (
|
||||||
|
pattern_start[1].get("type", AggregationType.SENTENCE)
|
||||||
|
if pattern_start
|
||||||
|
else AggregationType.SENTENCE
|
||||||
|
)
|
||||||
|
return Aggregation(text=stripped_text, type=type)
|
||||||
|
|
||||||
def add_pattern_pair(
|
def add_pattern(
|
||||||
self, pattern_id: str, start_pattern: str, end_pattern: str, remove_match: bool = True
|
self,
|
||||||
|
type: str,
|
||||||
|
start_pattern: str,
|
||||||
|
end_pattern: str,
|
||||||
|
action: MatchAction = MatchAction.REMOVE,
|
||||||
) -> "PatternPairAggregator":
|
) -> "PatternPairAggregator":
|
||||||
"""Add a pattern pair to detect in the text.
|
"""Add a pattern pair to detect in the text.
|
||||||
|
|
||||||
@@ -93,41 +131,94 @@ class PatternPairAggregator(BaseTextAggregator):
|
|||||||
the end pattern, and treat the content between them as a match.
|
the end pattern, and treat the content between them as a match.
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
pattern_id: Unique identifier for this pattern pair.
|
type: Identifier for this pattern pair. Should be unique and ideally descriptive.
|
||||||
|
(e.g., 'code', 'speaker', 'custom'). type can not be 'sentence' or 'word' as
|
||||||
|
those are reserved for the default behavior.
|
||||||
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.
|
action: What to do when a complete pattern is matched:
|
||||||
|
- MatchAction.REMOVE: Remove the matched pattern from the text.
|
||||||
|
- MatchAction.KEEP: Keep the matched pattern in the text and treat it as
|
||||||
|
normal text. This allows you to register handlers for
|
||||||
|
the pattern without affecting the aggregation logic.
|
||||||
|
- MatchAction.AGGREGATE: Return the matched pattern as a separate
|
||||||
|
aggregation object.
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
Self for method chaining.
|
Self for method chaining.
|
||||||
"""
|
"""
|
||||||
self._patterns[pattern_id] = {
|
if type in [AggregationType.SENTENCE, AggregationType.WORD]:
|
||||||
|
raise ValueError(
|
||||||
|
f"The aggregation type '{type}' is reserved for default behavior and can not be used for custom patterns."
|
||||||
|
)
|
||||||
|
self._patterns[type] = {
|
||||||
"start": start_pattern,
|
"start": start_pattern,
|
||||||
"end": end_pattern,
|
"end": end_pattern,
|
||||||
"remove_match": remove_match,
|
"type": type,
|
||||||
|
"action": action,
|
||||||
}
|
}
|
||||||
return self
|
return self
|
||||||
|
|
||||||
|
def add_pattern_pair(
|
||||||
|
self, pattern_id: str, start_pattern: str, end_pattern: str, remove_match: bool = True
|
||||||
|
):
|
||||||
|
"""Add a pattern pair to detect in the text.
|
||||||
|
|
||||||
|
.. deprecated:: 0.0.95
|
||||||
|
This function is deprecated and will be removed in a future version.
|
||||||
|
Use `add_pattern` with a type and MatchAction instead.
|
||||||
|
|
||||||
|
This method calls `add_pattern` setting type with the provided pattern_id and action
|
||||||
|
to either MatchAction.REMOVE or MatchAction.KEEP based on `remove_match`.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
pattern_id: Identifier for this pattern pair. Should be unique and ideally descriptive.
|
||||||
|
(e.g., 'code', 'speaker', 'custom'). pattern_id can not be 'sentence' or 'word'
|
||||||
|
as those arereserved for the default behavior.
|
||||||
|
start_pattern: Pattern that marks the beginning of content.
|
||||||
|
end_pattern: Pattern that marks the end of content.
|
||||||
|
remove_match: If True, the matched pattern will be removed from the text. (Same as MatchAction.REMOVE)
|
||||||
|
If False, it will be kept and treated as normal text. (Same as MatchAction.KEEP)
|
||||||
|
"""
|
||||||
|
import warnings
|
||||||
|
|
||||||
|
with warnings.catch_warnings():
|
||||||
|
warnings.simplefilter("once")
|
||||||
|
warnings.warn(
|
||||||
|
"add_pattern_pair with a pattern_id or remove_match is deprecated and will be"
|
||||||
|
" removed in a future version. Use add_pattern with a type and MatchAction instead",
|
||||||
|
DeprecationWarning,
|
||||||
|
stacklevel=2,
|
||||||
|
)
|
||||||
|
|
||||||
|
action = MatchAction.REMOVE if remove_match else MatchAction.KEEP
|
||||||
|
return self.add_pattern(
|
||||||
|
type=pattern_id,
|
||||||
|
start_pattern=start_pattern,
|
||||||
|
end_pattern=end_pattern,
|
||||||
|
action=action,
|
||||||
|
)
|
||||||
|
|
||||||
def on_pattern_match(
|
def on_pattern_match(
|
||||||
self, pattern_id: str, handler: Callable[[PatternMatch], Awaitable[None]]
|
self, type: str, handler: Callable[[PatternMatch], Awaitable[None]]
|
||||||
) -> "PatternPairAggregator":
|
) -> "PatternPairAggregator":
|
||||||
"""Register a handler for when a pattern pair is matched.
|
"""Register a handler for when a pattern pair is matched.
|
||||||
|
|
||||||
The handler will be called whenever a complete match for the
|
The handler will be called whenever a complete match for the
|
||||||
specified pattern ID is found in the text.
|
specified type is found in the text.
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
pattern_id: ID of the pattern pair to match.
|
type: The type of the pattern pair to trigger the handler.
|
||||||
handler: Async function to call when pattern is matched.
|
handler: Async function to call when pattern is matched.
|
||||||
The function should accept a PatternMatch object.
|
The function should accept a PatternMatch object.
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
Self for method chaining.
|
Self for method chaining.
|
||||||
"""
|
"""
|
||||||
self._handlers[pattern_id] = handler
|
self._handlers[type] = 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 +228,19 @@ 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 type, 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"]
|
action = pattern_info["action"]
|
||||||
|
|
||||||
# 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,24 +256,25 @@ 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
|
content=content.strip(), type=type, full_match=full_match
|
||||||
)
|
)
|
||||||
|
|
||||||
# Call the appropriate handler if registered
|
# Call the appropriate handler if registered
|
||||||
if pattern_id in self._handlers:
|
if type in self._handlers:
|
||||||
try:
|
try:
|
||||||
await self._handlers[pattern_id](pattern_match)
|
await self._handlers[type](pattern_match)
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.error(f"Error in pattern handler for {pattern_id}: {e}")
|
logger.error(f"Error in pattern handler for {type}: {e}")
|
||||||
|
|
||||||
# Remove the pattern from the text if configured
|
# Remove the pattern from the text if configured
|
||||||
if remove_match:
|
if action == MatchAction.REMOVE:
|
||||||
processed_text = processed_text.replace(full_match, "", 1)
|
processed_text = processed_text.replace(full_match, "", 1)
|
||||||
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, dict]]:
|
||||||
"""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,9 +284,10 @@ 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, pattern_info) 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 type, pattern_info in self._patterns.items():
|
||||||
start = pattern_info["start"]
|
start = pattern_info["start"]
|
||||||
end = pattern_info["end"]
|
end = pattern_info["end"]
|
||||||
|
|
||||||
@@ -203,12 +296,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]
|
||||||
|
|
||||||
return False
|
return None
|
||||||
|
|
||||||
async def aggregate(self, text: str) -> Optional[Aggregation]:
|
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
|
||||||
@@ -220,24 +317,43 @@ class PatternPairAggregator(BaseTextAggregator):
|
|||||||
text: New text to add to the buffer.
|
text: New text to add to the buffer.
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
An Aggregation object containing processed text up to a sentence boundary
|
Processed text up to a sentence boundary, or None if more
|
||||||
and marked as SENTENCE type, or None if more text is needed to form a
|
text is needed to form a complete sentence or pattern.
|
||||||
complete sentence or pattern.
|
|
||||||
"""
|
"""
|
||||||
# Add new text to buffer
|
# Add new text to buffer
|
||||||
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.type for p in patterns]}. Only the first pattern will be returned."
|
||||||
|
)
|
||||||
|
# If the pattern found is set to be aggregated, return it
|
||||||
|
action = self._patterns[patterns[0].type].get("action", MatchAction.REMOVE)
|
||||||
|
if action == MatchAction.AGGREGATE:
|
||||||
|
self._text = ""
|
||||||
|
return patterns[0]
|
||||||
|
|
||||||
# Check if we have incomplete patterns
|
# Check if we have incomplete patterns
|
||||||
if self._has_incomplete_patterns(self._text):
|
pattern_start = self._match_start_of_pattern(self._text)
|
||||||
# Still waiting for complete patterns
|
if pattern_start is not None:
|
||||||
return None
|
# If the start pattern is at the beginning or should not be separately aggregated, return None
|
||||||
|
if (
|
||||||
|
pattern_start[0] == 0
|
||||||
|
or pattern_start[1].get("action", MatchAction.REMOVE) != MatchAction.AGGREGATE
|
||||||
|
):
|
||||||
|
return None
|
||||||
|
# Otherwise, strip the text up to the start pattern and return it
|
||||||
|
result = self._text[: pattern_start[0]]
|
||||||
|
self._text = self._text[pattern_start[0] :]
|
||||||
|
return PatternMatch(
|
||||||
|
content=result.strip(), type=AggregationType.SENTENCE, full_match=result
|
||||||
|
)
|
||||||
|
|
||||||
# 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)
|
||||||
@@ -245,7 +361,9 @@ 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 Aggregation(text=result.strip(), type=AggregationType.SENTENCE)
|
return PatternMatch(
|
||||||
|
content=result.strip(), type=AggregationType.SENTENCE, full_match=result
|
||||||
|
)
|
||||||
|
|
||||||
# No complete sentence found yet
|
# No complete sentence found yet
|
||||||
return None
|
return None
|
||||||
|
|||||||
@@ -7,30 +7,42 @@
|
|||||||
import unittest
|
import unittest
|
||||||
from unittest.mock import AsyncMock
|
from unittest.mock import AsyncMock
|
||||||
|
|
||||||
from pipecat.utils.text.pattern_pair_aggregator import PatternMatch, PatternPairAggregator
|
from pipecat.utils.text.pattern_pair_aggregator import (
|
||||||
|
MatchAction,
|
||||||
|
PatternMatch,
|
||||||
|
PatternPairAggregator,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
class TestPatternPairAggregator(unittest.IsolatedAsyncioTestCase):
|
class TestPatternPairAggregator(unittest.IsolatedAsyncioTestCase):
|
||||||
def setUp(self):
|
def setUp(self):
|
||||||
self.aggregator = PatternPairAggregator()
|
self.aggregator = PatternPairAggregator()
|
||||||
self.test_handler = AsyncMock()
|
self.test_handler = AsyncMock()
|
||||||
|
self.code_handler = AsyncMock()
|
||||||
|
|
||||||
# Add a test pattern
|
# Add a test pattern
|
||||||
self.aggregator.add_pattern_pair(
|
self.aggregator.add_pattern_pair(
|
||||||
pattern_id="test_pattern",
|
pattern_id="test_pattern",
|
||||||
start_pattern="<test>",
|
start_pattern="<test>",
|
||||||
end_pattern="</test>",
|
end_pattern="</test>",
|
||||||
remove_match=True,
|
)
|
||||||
|
self.aggregator.add_pattern(
|
||||||
|
type="code_pattern",
|
||||||
|
start_pattern="<code>",
|
||||||
|
end_pattern="</code>",
|
||||||
|
action=MatchAction.AGGREGATE,
|
||||||
)
|
)
|
||||||
|
|
||||||
# Register the mock handler
|
# Register the mock handler
|
||||||
self.aggregator.on_pattern_match("test_pattern", self.test_handler)
|
self.aggregator.on_pattern_match("test_pattern", self.test_handler)
|
||||||
|
self.aggregator.on_pattern_match("code_pattern", self.code_handler)
|
||||||
|
|
||||||
async def test_pattern_match_and_removal(self):
|
async def test_pattern_match_and_removal(self):
|
||||||
# First part doesn't complete the pattern
|
# First part doesn't complete the pattern
|
||||||
result = await self.aggregator.aggregate("Hello <test>pattern")
|
result = await self.aggregator.aggregate("Hello <test>pattern")
|
||||||
self.assertIsNone(result)
|
self.assertIsNone(result)
|
||||||
self.assertEqual(self.aggregator.text.text, "Hello <test>pattern")
|
self.assertEqual(self.aggregator.text.text, "Hello <test>pattern")
|
||||||
|
self.assertEqual(self.aggregator.text.type, "test_pattern")
|
||||||
|
|
||||||
# Second part completes the pattern and includes an exclamation point
|
# Second part completes the pattern and includes an exclamation point
|
||||||
result = await self.aggregator.aggregate(" content</test>!")
|
result = await self.aggregator.aggregate(" content</test>!")
|
||||||
@@ -39,9 +51,9 @@ class TestPatternPairAggregator(unittest.IsolatedAsyncioTestCase):
|
|||||||
self.test_handler.assert_called_once()
|
self.test_handler.assert_called_once()
|
||||||
call_args = self.test_handler.call_args[0][0]
|
call_args = self.test_handler.call_args[0][0]
|
||||||
self.assertIsInstance(call_args, PatternMatch)
|
self.assertIsInstance(call_args, PatternMatch)
|
||||||
self.assertEqual(call_args.pattern_id, "test_pattern")
|
self.assertEqual(call_args.type, "test_pattern")
|
||||||
self.assertEqual(call_args.full_match, "<test>pattern content</test>")
|
self.assertEqual(call_args.full_match, "<test>pattern content</test>")
|
||||||
self.assertEqual(call_args.content, "pattern content")
|
self.assertEqual(call_args.text, "pattern content")
|
||||||
|
|
||||||
# The exclamation point should be treated as a sentence boundary,
|
# The exclamation point should be treated as a sentence boundary,
|
||||||
# so the result should include just text up to and including "!"
|
# so the result should include just text up to and including "!"
|
||||||
@@ -52,7 +64,35 @@ class TestPatternPairAggregator(unittest.IsolatedAsyncioTestCase):
|
|||||||
# should be stripped in the returned Aggregation.
|
# should be stripped in the returned Aggregation.
|
||||||
result = await self.aggregator.aggregate(" This is another sentence.")
|
result = await self.aggregator.aggregate(" This is another sentence.")
|
||||||
self.assertEqual(result.text, "This is another sentence.")
|
self.assertEqual(result.text, "This is another sentence.")
|
||||||
|
|
||||||
|
# Buffer should be empty after returning a complete sentence
|
||||||
|
self.assertEqual(self.aggregator.text.text, "")
|
||||||
|
|
||||||
|
async def test_pattern_match_and_aggregate(self):
|
||||||
|
# First part doesn't complete the pattern
|
||||||
|
result = await self.aggregator.aggregate("Here is code <code>pattern")
|
||||||
|
self.assertEqual(result.text, "Here is code")
|
||||||
|
self.assertEqual(self.aggregator.text.text, "<code>pattern")
|
||||||
|
self.assertEqual(self.aggregator.text.type, "code_pattern")
|
||||||
|
|
||||||
|
# Second part completes the pattern and includes an exclamation point
|
||||||
|
result = await self.aggregator.aggregate(" content</code>")
|
||||||
|
|
||||||
|
# Verify the handler was called with correct PatternMatch object
|
||||||
|
self.code_handler.assert_called_once()
|
||||||
|
call_args = self.code_handler.call_args[0][0]
|
||||||
|
self.assertIsInstance(call_args, PatternMatch)
|
||||||
|
self.assertEqual(call_args.type, "code_pattern")
|
||||||
|
self.assertEqual(call_args.full_match, "<code>pattern content</code>")
|
||||||
|
self.assertEqual(call_args.text, "pattern content")
|
||||||
|
self.assertEqual(result.text, "pattern content")
|
||||||
|
self.assertEqual(result.type, "code_pattern")
|
||||||
|
|
||||||
|
# Next sentence should be processed separately
|
||||||
|
result = await self.aggregator.aggregate(" This is another sentence.")
|
||||||
|
self.assertEqual(result.text, "This is another sentence.")
|
||||||
self.assertEqual(result.type, "sentence")
|
self.assertEqual(result.type, "sentence")
|
||||||
|
|
||||||
# Buffer should be empty after returning a complete sentence
|
# Buffer should be empty after returning a complete sentence
|
||||||
self.assertEqual(self.aggregator.text.text, "")
|
self.assertEqual(self.aggregator.text.text, "")
|
||||||
|
|
||||||
@@ -68,6 +108,7 @@ class TestPatternPairAggregator(unittest.IsolatedAsyncioTestCase):
|
|||||||
|
|
||||||
# Buffer should contain the incomplete text
|
# Buffer should contain the incomplete text
|
||||||
self.assertEqual(self.aggregator.text.text, "Hello <test>pattern content")
|
self.assertEqual(self.aggregator.text.text, "Hello <test>pattern content")
|
||||||
|
self.assertEqual(self.aggregator.text.type, "test_pattern")
|
||||||
|
|
||||||
# Reset and confirm buffer is cleared
|
# Reset and confirm buffer is cleared
|
||||||
await self.aggregator.reset()
|
await self.aggregator.reset()
|
||||||
@@ -78,15 +119,18 @@ class TestPatternPairAggregator(unittest.IsolatedAsyncioTestCase):
|
|||||||
voice_handler = AsyncMock()
|
voice_handler = AsyncMock()
|
||||||
emphasis_handler = AsyncMock()
|
emphasis_handler = AsyncMock()
|
||||||
|
|
||||||
self.aggregator.add_pattern_pair(
|
self.aggregator.add_pattern(
|
||||||
pattern_id="voice", start_pattern="<voice>", end_pattern="</voice>", remove_match=True
|
type="voice",
|
||||||
|
start_pattern="<voice>",
|
||||||
|
end_pattern="</voice>",
|
||||||
|
action=MatchAction.REMOVE,
|
||||||
)
|
)
|
||||||
|
|
||||||
self.aggregator.add_pattern_pair(
|
self.aggregator.add_pattern(
|
||||||
pattern_id="emphasis",
|
type="emphasis",
|
||||||
start_pattern="<em>",
|
start_pattern="<em>",
|
||||||
end_pattern="</em>",
|
end_pattern="</em>",
|
||||||
remove_match=False, # Keep emphasis tags
|
action=MatchAction.KEEP, # Keep emphasis tags
|
||||||
)
|
)
|
||||||
|
|
||||||
self.aggregator.on_pattern_match("voice", voice_handler)
|
self.aggregator.on_pattern_match("voice", voice_handler)
|
||||||
@@ -99,17 +143,16 @@ class TestPatternPairAggregator(unittest.IsolatedAsyncioTestCase):
|
|||||||
# Both handlers should be called with correct data
|
# Both handlers should be called with correct data
|
||||||
voice_handler.assert_called_once()
|
voice_handler.assert_called_once()
|
||||||
voice_match = voice_handler.call_args[0][0]
|
voice_match = voice_handler.call_args[0][0]
|
||||||
self.assertEqual(voice_match.pattern_id, "voice")
|
self.assertEqual(voice_match.type, "voice")
|
||||||
self.assertEqual(voice_match.content, "female")
|
self.assertEqual(voice_match.text, "female")
|
||||||
|
|
||||||
emphasis_handler.assert_called_once()
|
emphasis_handler.assert_called_once()
|
||||||
emphasis_match = emphasis_handler.call_args[0][0]
|
emphasis_match = emphasis_handler.call_args[0][0]
|
||||||
self.assertEqual(emphasis_match.pattern_id, "emphasis")
|
self.assertEqual(emphasis_match.type, "emphasis")
|
||||||
self.assertEqual(emphasis_match.content, "very")
|
self.assertEqual(emphasis_match.text, "very")
|
||||||
|
|
||||||
# Voice pattern should be removed, emphasis pattern should remain
|
# Voice pattern should be removed, emphasis pattern should remain
|
||||||
self.assertEqual(result.text, "Hello I am <em>very</em> excited to meet you!")
|
self.assertEqual(result.text, "Hello I am <em>very</em> excited to meet you!")
|
||||||
self.assertEqual(result.type, "sentence")
|
|
||||||
|
|
||||||
# Buffer should be empty
|
# Buffer should be empty
|
||||||
self.assertEqual(self.aggregator.text.text, "")
|
self.assertEqual(self.aggregator.text.text, "")
|
||||||
@@ -141,11 +184,10 @@ class TestPatternPairAggregator(unittest.IsolatedAsyncioTestCase):
|
|||||||
# Handler should be called with entire content
|
# Handler should be called with entire content
|
||||||
self.test_handler.assert_called_once()
|
self.test_handler.assert_called_once()
|
||||||
call_args = self.test_handler.call_args[0][0]
|
call_args = self.test_handler.call_args[0][0]
|
||||||
self.assertEqual(call_args.content, "This is sentence one. This is sentence two.")
|
self.assertEqual(call_args.text, "This is sentence one. This is sentence two.")
|
||||||
|
|
||||||
# Pattern should be removed, resulting in text with sentences merged
|
# Pattern should be removed, resulting in text with sentences merged
|
||||||
self.assertEqual(result.text, "Hello Final sentence.")
|
self.assertEqual(result.text, "Hello Final sentence.")
|
||||||
self.assertEqual(result.type, "sentence")
|
|
||||||
|
|
||||||
# Buffer should be empty
|
# Buffer should be empty
|
||||||
self.assertEqual(self.aggregator.text.text, "")
|
self.assertEqual(self.aggregator.text.text, "")
|
||||||
|
|||||||
Reference in New Issue
Block a user