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:
mattie ruth backman
2025-11-14 18:31:58 -05:00
committed by Mattie Ruth
parent dcc20f86e1
commit 24266c238f
5 changed files with 283 additions and 88 deletions

View File

@@ -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

View File

@@ -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"))

View File

@@ -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()

View File

@@ -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
if modified:
self._text = processed_text 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:
# 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 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

View File

@@ -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, "")