Add tests, improve frame timing
This commit is contained in:
@@ -45,60 +45,81 @@ class DTMFAggregator(FrameProcessor):
|
|||||||
**kwargs,
|
**kwargs,
|
||||||
):
|
):
|
||||||
super().__init__(**kwargs)
|
super().__init__(**kwargs)
|
||||||
|
|
||||||
self._aggregation = ""
|
self._aggregation = ""
|
||||||
self._idle_timeout = timeout
|
self._idle_timeout = timeout
|
||||||
self._termination_digit = termination_digit
|
self._termination_digit = termination_digit
|
||||||
self._prefix = prefix
|
self._prefix = prefix
|
||||||
|
|
||||||
self._digit_event = asyncio.Event()
|
self._digit_event = asyncio.Event()
|
||||||
self._aggregation_task = None
|
self._aggregation_task: Optional[asyncio.Task] = None
|
||||||
|
|
||||||
|
def _create_aggregation_task(self) -> None:
|
||||||
|
"""Creates the aggregation task if it hasn't been created yet."""
|
||||||
|
if not self._aggregation_task:
|
||||||
|
self._aggregation_task = self.create_task(self._aggregation_task_handler())
|
||||||
|
|
||||||
|
async def _stop(self) -> None:
|
||||||
|
"""Stops and cleans up the aggregation task."""
|
||||||
|
if self._aggregation_task:
|
||||||
|
await self.cancel_task(self._aggregation_task)
|
||||||
|
self._aggregation_task = None
|
||||||
|
|
||||||
|
async def cleanup(self) -> None:
|
||||||
|
await super().cleanup()
|
||||||
|
if self._aggregation_task:
|
||||||
|
await self._stop()
|
||||||
|
|
||||||
async def process_frame(self, frame: Frame, direction: FrameDirection) -> None:
|
async def process_frame(self, frame: Frame, direction: FrameDirection) -> None:
|
||||||
await super().process_frame(frame, direction)
|
await super().process_frame(frame, direction)
|
||||||
|
|
||||||
if isinstance(frame, InputDTMFFrame):
|
if isinstance(frame, (EndFrame, CancelFrame)):
|
||||||
await self._handle_dtmf_frame(frame, direction)
|
# Flush any pending aggregation before stopping
|
||||||
|
if self._aggregation:
|
||||||
|
await self._flush_aggregation()
|
||||||
|
await self._stop()
|
||||||
|
await self.push_frame(frame, direction)
|
||||||
|
return
|
||||||
elif isinstance(frame, StartInterruptionFrame):
|
elif isinstance(frame, StartInterruptionFrame):
|
||||||
# Flush on interruption
|
|
||||||
if self._aggregation:
|
if self._aggregation:
|
||||||
await self._flush_aggregation(direction)
|
await self._flush_aggregation()
|
||||||
elif isinstance(frame, (EndFrame, CancelFrame)):
|
await self.push_frame(frame, direction)
|
||||||
# Flush any pending aggregation
|
return
|
||||||
if self._aggregation:
|
|
||||||
await self._flush_aggregation(direction)
|
|
||||||
await self._stop_aggregation_task()
|
|
||||||
|
|
||||||
# Push frame
|
|
||||||
await self.push_frame(frame, direction)
|
await self.push_frame(frame, direction)
|
||||||
|
|
||||||
async def _handle_dtmf_frame(self, frame: InputDTMFFrame, direction: FrameDirection):
|
if isinstance(frame, InputDTMFFrame):
|
||||||
# Start aggregation task if not running
|
await self._handle_dtmf_frame(frame)
|
||||||
if self._aggregation_task is None:
|
|
||||||
self._aggregation_task = self.create_task(self._aggregation_task_handler(direction))
|
async def _handle_dtmf_frame(
|
||||||
|
self,
|
||||||
|
frame: InputDTMFFrame,
|
||||||
|
):
|
||||||
|
# Create task on first DTMF input if needed
|
||||||
|
if not self._aggregation_task:
|
||||||
|
self._create_aggregation_task()
|
||||||
|
|
||||||
# Add digit to aggregation
|
|
||||||
digit_value = frame.button.value
|
digit_value = frame.button.value
|
||||||
self._aggregation += digit_value
|
self._aggregation += digit_value
|
||||||
|
|
||||||
# Check for immediate flush conditions
|
# Check for immediate flush conditions
|
||||||
if frame.button == self._termination_digit:
|
if frame.button == self._termination_digit:
|
||||||
await self._flush_aggregation(direction)
|
await self._flush_aggregation()
|
||||||
else:
|
else:
|
||||||
# Signal new digit received
|
# Signal new digit received
|
||||||
self._digit_event.set()
|
self._digit_event.set()
|
||||||
|
|
||||||
async def _aggregation_task_handler(self, direction: FrameDirection):
|
async def _aggregation_task_handler(self):
|
||||||
"""Background task that handles timeout-based flushing."""
|
"""Background task that handles timeout-based flushing."""
|
||||||
while True:
|
while True:
|
||||||
try:
|
try:
|
||||||
await asyncio.wait_for(self._digit_event.wait(), timeout=self._idle_timeout)
|
await asyncio.wait_for(self._digit_event.wait(), timeout=self._idle_timeout)
|
||||||
self._digit_event.clear()
|
|
||||||
except asyncio.TimeoutError:
|
except asyncio.TimeoutError:
|
||||||
if self._aggregation:
|
if self._aggregation:
|
||||||
await self._flush_aggregation(direction)
|
await self._flush_aggregation()
|
||||||
|
finally:
|
||||||
|
self._digit_event.clear()
|
||||||
|
|
||||||
async def _flush_aggregation(self, direction: FrameDirection):
|
async def _flush_aggregation(self):
|
||||||
"""Flush the current aggregation as a TranscriptionFrame."""
|
"""Flush the current aggregation as a TranscriptionFrame."""
|
||||||
if not self._aggregation:
|
if not self._aggregation:
|
||||||
return
|
return
|
||||||
@@ -113,16 +134,7 @@ class DTMFAggregator(FrameProcessor):
|
|||||||
text=transcription_text, user_id="", timestamp=time_now_iso8601()
|
text=transcription_text, user_id="", timestamp=time_now_iso8601()
|
||||||
)
|
)
|
||||||
|
|
||||||
await self.push_frame(transcription_frame, direction)
|
await self.push_frame(transcription_frame)
|
||||||
|
|
||||||
# Reset aggregation
|
# Reset aggregation
|
||||||
self._aggregation = ""
|
self._aggregation = ""
|
||||||
|
|
||||||
async def _stop_aggregation_task(self):
|
|
||||||
if self._aggregation_task:
|
|
||||||
await self.cancel_task(self._aggregation_task)
|
|
||||||
self._aggregation_task = None
|
|
||||||
|
|
||||||
async def cleanup(self) -> None:
|
|
||||||
await super().cleanup()
|
|
||||||
await self._stop_aggregation_task()
|
|
||||||
|
|||||||
257
tests/test_dtmf_aggregator.py
Normal file
257
tests/test_dtmf_aggregator.py
Normal file
@@ -0,0 +1,257 @@
|
|||||||
|
#
|
||||||
|
# Copyright (c) 2024-2025 Daily
|
||||||
|
#
|
||||||
|
# SPDX-License-Identifier: BSD 2-Clause License
|
||||||
|
#
|
||||||
|
|
||||||
|
import unittest
|
||||||
|
|
||||||
|
from pipecat.frames.frames import (
|
||||||
|
EndFrame,
|
||||||
|
InputDTMFFrame,
|
||||||
|
KeypadEntry,
|
||||||
|
StartInterruptionFrame,
|
||||||
|
TranscriptionFrame,
|
||||||
|
)
|
||||||
|
from pipecat.processors.aggregators.dtmf_aggregator import DTMFAggregator
|
||||||
|
from pipecat.tests.utils import SleepFrame, run_test
|
||||||
|
|
||||||
|
|
||||||
|
class TestDTMFAggregator(unittest.IsolatedAsyncioTestCase):
|
||||||
|
async def test_basic_aggregation_with_pound(self):
|
||||||
|
"""Test basic DTMF aggregation ending with pound key."""
|
||||||
|
aggregator = DTMFAggregator()
|
||||||
|
frames_to_send = [
|
||||||
|
InputDTMFFrame(button=KeypadEntry.ONE),
|
||||||
|
InputDTMFFrame(button=KeypadEntry.TWO),
|
||||||
|
InputDTMFFrame(button=KeypadEntry.THREE),
|
||||||
|
InputDTMFFrame(button=KeypadEntry.POUND),
|
||||||
|
]
|
||||||
|
expected_down_frames = [
|
||||||
|
InputDTMFFrame,
|
||||||
|
InputDTMFFrame,
|
||||||
|
InputDTMFFrame,
|
||||||
|
InputDTMFFrame,
|
||||||
|
TranscriptionFrame,
|
||||||
|
]
|
||||||
|
|
||||||
|
received_down_frames, _ = await run_test(
|
||||||
|
aggregator,
|
||||||
|
frames_to_send=frames_to_send,
|
||||||
|
expected_down_frames=expected_down_frames,
|
||||||
|
)
|
||||||
|
|
||||||
|
# Find and verify the TranscriptionFrame
|
||||||
|
transcription_frames = [
|
||||||
|
f for f in received_down_frames if isinstance(f, TranscriptionFrame)
|
||||||
|
]
|
||||||
|
self.assertEqual(len(transcription_frames), 1)
|
||||||
|
self.assertEqual(transcription_frames[0].text, "DTMF: 123#")
|
||||||
|
|
||||||
|
async def test_timeout_aggregation(self):
|
||||||
|
"""Test DTMF aggregation with timeout flush."""
|
||||||
|
aggregator = DTMFAggregator(timeout=0.1)
|
||||||
|
frames_to_send = [
|
||||||
|
InputDTMFFrame(button=KeypadEntry.ONE),
|
||||||
|
InputDTMFFrame(button=KeypadEntry.TWO),
|
||||||
|
SleepFrame(sleep=0.2), # This should trigger timeout
|
||||||
|
InputDTMFFrame(button=KeypadEntry.THREE),
|
||||||
|
SleepFrame(sleep=0.2), # This should trigger another timeout
|
||||||
|
]
|
||||||
|
expected_down_frames = [
|
||||||
|
InputDTMFFrame,
|
||||||
|
InputDTMFFrame,
|
||||||
|
TranscriptionFrame, # First aggregation "12"
|
||||||
|
InputDTMFFrame,
|
||||||
|
TranscriptionFrame, # Second aggregation "3"
|
||||||
|
]
|
||||||
|
|
||||||
|
received_down_frames, _ = await run_test(
|
||||||
|
aggregator,
|
||||||
|
frames_to_send=frames_to_send,
|
||||||
|
expected_down_frames=expected_down_frames,
|
||||||
|
)
|
||||||
|
|
||||||
|
# Find the TranscriptionFrames
|
||||||
|
transcription_frames = [
|
||||||
|
f for f in received_down_frames if isinstance(f, TranscriptionFrame)
|
||||||
|
]
|
||||||
|
self.assertEqual(len(transcription_frames), 2)
|
||||||
|
self.assertEqual(transcription_frames[0].text, "DTMF: 12")
|
||||||
|
self.assertEqual(transcription_frames[1].text, "DTMF: 3")
|
||||||
|
|
||||||
|
async def test_multiple_aggregations(self):
|
||||||
|
"""Test multiple DTMF sequences with pound termination."""
|
||||||
|
aggregator = DTMFAggregator(timeout=0.1)
|
||||||
|
frames_to_send = [
|
||||||
|
InputDTMFFrame(button=KeypadEntry.ONE),
|
||||||
|
InputDTMFFrame(button=KeypadEntry.TWO),
|
||||||
|
InputDTMFFrame(button=KeypadEntry.POUND), # First sequence
|
||||||
|
InputDTMFFrame(button=KeypadEntry.FOUR),
|
||||||
|
InputDTMFFrame(button=KeypadEntry.FIVE),
|
||||||
|
SleepFrame(sleep=0.2), # Second sequence via timeout
|
||||||
|
]
|
||||||
|
expected_down_frames = [
|
||||||
|
InputDTMFFrame,
|
||||||
|
InputDTMFFrame,
|
||||||
|
InputDTMFFrame,
|
||||||
|
TranscriptionFrame, # "12#"
|
||||||
|
InputDTMFFrame,
|
||||||
|
InputDTMFFrame,
|
||||||
|
TranscriptionFrame, # "45"
|
||||||
|
]
|
||||||
|
|
||||||
|
received_down_frames, _ = await run_test(
|
||||||
|
aggregator,
|
||||||
|
frames_to_send=frames_to_send,
|
||||||
|
expected_down_frames=expected_down_frames,
|
||||||
|
)
|
||||||
|
|
||||||
|
transcription_frames = [
|
||||||
|
f for f in received_down_frames if isinstance(f, TranscriptionFrame)
|
||||||
|
]
|
||||||
|
self.assertEqual(len(transcription_frames), 2)
|
||||||
|
self.assertEqual(transcription_frames[0].text, "DTMF: 12#")
|
||||||
|
self.assertEqual(transcription_frames[1].text, "DTMF: 45")
|
||||||
|
|
||||||
|
async def test_end_frame_flush(self):
|
||||||
|
"""Test that EndFrame flushes pending aggregation."""
|
||||||
|
aggregator = DTMFAggregator(timeout=1.0)
|
||||||
|
frames_to_send = [
|
||||||
|
InputDTMFFrame(button=KeypadEntry.ONE),
|
||||||
|
InputDTMFFrame(button=KeypadEntry.TWO),
|
||||||
|
SleepFrame(sleep=0.1), # Allow time for aggregation
|
||||||
|
EndFrame(),
|
||||||
|
]
|
||||||
|
expected_down_frames = [
|
||||||
|
InputDTMFFrame,
|
||||||
|
InputDTMFFrame,
|
||||||
|
TranscriptionFrame, # Should flush before EndFrame
|
||||||
|
EndFrame,
|
||||||
|
]
|
||||||
|
|
||||||
|
received_down_frames, _ = await run_test(
|
||||||
|
aggregator,
|
||||||
|
frames_to_send=frames_to_send,
|
||||||
|
expected_down_frames=expected_down_frames,
|
||||||
|
send_end_frame=False, # We're sending one in the test to test EndFrame logic
|
||||||
|
)
|
||||||
|
|
||||||
|
transcription_frames = [
|
||||||
|
f for f in received_down_frames if isinstance(f, TranscriptionFrame)
|
||||||
|
]
|
||||||
|
self.assertEqual(len(transcription_frames), 1)
|
||||||
|
self.assertEqual(transcription_frames[0].text, "DTMF: 12")
|
||||||
|
|
||||||
|
async def test_interruption_frame_flush(self):
|
||||||
|
"""Test that StartInterruptionFrame flushes pending aggregation."""
|
||||||
|
aggregator = DTMFAggregator(timeout=1.0)
|
||||||
|
frames_to_send = [
|
||||||
|
InputDTMFFrame(button=KeypadEntry.ONE),
|
||||||
|
InputDTMFFrame(button=KeypadEntry.TWO),
|
||||||
|
SleepFrame(sleep=0.1), # Allow time for aggregation
|
||||||
|
StartInterruptionFrame(),
|
||||||
|
]
|
||||||
|
expected_down_frames = [
|
||||||
|
InputDTMFFrame,
|
||||||
|
InputDTMFFrame,
|
||||||
|
TranscriptionFrame, # Should flush before interruption
|
||||||
|
StartInterruptionFrame,
|
||||||
|
]
|
||||||
|
|
||||||
|
received_down_frames, _ = await run_test(
|
||||||
|
aggregator,
|
||||||
|
frames_to_send=frames_to_send,
|
||||||
|
expected_down_frames=expected_down_frames,
|
||||||
|
)
|
||||||
|
|
||||||
|
transcription_frames = [
|
||||||
|
f for f in received_down_frames if isinstance(f, TranscriptionFrame)
|
||||||
|
]
|
||||||
|
self.assertEqual(len(transcription_frames), 1)
|
||||||
|
self.assertEqual(transcription_frames[0].text, "DTMF: 12")
|
||||||
|
|
||||||
|
async def test_custom_prefix(self):
|
||||||
|
"""Test custom prefix configuration."""
|
||||||
|
aggregator = DTMFAggregator(prefix="Menu: ")
|
||||||
|
frames_to_send = [
|
||||||
|
InputDTMFFrame(button=KeypadEntry.ONE),
|
||||||
|
InputDTMFFrame(button=KeypadEntry.POUND),
|
||||||
|
]
|
||||||
|
expected_down_frames = [
|
||||||
|
InputDTMFFrame,
|
||||||
|
InputDTMFFrame,
|
||||||
|
TranscriptionFrame,
|
||||||
|
]
|
||||||
|
|
||||||
|
received_down_frames, _ = await run_test(
|
||||||
|
aggregator,
|
||||||
|
frames_to_send=frames_to_send,
|
||||||
|
expected_down_frames=expected_down_frames,
|
||||||
|
)
|
||||||
|
|
||||||
|
transcription_frames = [
|
||||||
|
f for f in received_down_frames if isinstance(f, TranscriptionFrame)
|
||||||
|
]
|
||||||
|
self.assertEqual(len(transcription_frames), 1)
|
||||||
|
self.assertEqual(transcription_frames[0].text, "Menu: 1#")
|
||||||
|
|
||||||
|
async def test_custom_termination_digit(self):
|
||||||
|
"""Test custom termination digit configuration."""
|
||||||
|
aggregator = DTMFAggregator(termination_digit=KeypadEntry.STAR)
|
||||||
|
frames_to_send = [
|
||||||
|
InputDTMFFrame(button=KeypadEntry.ONE),
|
||||||
|
InputDTMFFrame(button=KeypadEntry.TWO),
|
||||||
|
InputDTMFFrame(button=KeypadEntry.STAR), # Custom terminator
|
||||||
|
]
|
||||||
|
expected_down_frames = [
|
||||||
|
InputDTMFFrame,
|
||||||
|
InputDTMFFrame,
|
||||||
|
InputDTMFFrame,
|
||||||
|
TranscriptionFrame,
|
||||||
|
]
|
||||||
|
|
||||||
|
received_down_frames, _ = await run_test(
|
||||||
|
aggregator,
|
||||||
|
frames_to_send=frames_to_send,
|
||||||
|
expected_down_frames=expected_down_frames,
|
||||||
|
)
|
||||||
|
|
||||||
|
transcription_frames = [
|
||||||
|
f for f in received_down_frames if isinstance(f, TranscriptionFrame)
|
||||||
|
]
|
||||||
|
self.assertEqual(len(transcription_frames), 1)
|
||||||
|
self.assertEqual(transcription_frames[0].text, "DTMF: 12*")
|
||||||
|
|
||||||
|
async def test_all_keypad_entries(self):
|
||||||
|
"""Test all possible keypad entries."""
|
||||||
|
aggregator = DTMFAggregator()
|
||||||
|
frames_to_send = [
|
||||||
|
InputDTMFFrame(button=KeypadEntry.ZERO),
|
||||||
|
InputDTMFFrame(button=KeypadEntry.ONE),
|
||||||
|
InputDTMFFrame(button=KeypadEntry.TWO),
|
||||||
|
InputDTMFFrame(button=KeypadEntry.THREE),
|
||||||
|
InputDTMFFrame(button=KeypadEntry.FOUR),
|
||||||
|
InputDTMFFrame(button=KeypadEntry.FIVE),
|
||||||
|
InputDTMFFrame(button=KeypadEntry.SIX),
|
||||||
|
InputDTMFFrame(button=KeypadEntry.SEVEN),
|
||||||
|
InputDTMFFrame(button=KeypadEntry.EIGHT),
|
||||||
|
InputDTMFFrame(button=KeypadEntry.NINE),
|
||||||
|
InputDTMFFrame(button=KeypadEntry.STAR),
|
||||||
|
InputDTMFFrame(button=KeypadEntry.POUND),
|
||||||
|
]
|
||||||
|
|
||||||
|
# All the InputDTMFFrames plus one TranscriptionFrame
|
||||||
|
expected_down_frames = [InputDTMFFrame] * 12 + [TranscriptionFrame]
|
||||||
|
|
||||||
|
received_down_frames, _ = await run_test(
|
||||||
|
aggregator,
|
||||||
|
frames_to_send=frames_to_send,
|
||||||
|
expected_down_frames=expected_down_frames,
|
||||||
|
)
|
||||||
|
|
||||||
|
transcription_frames = [
|
||||||
|
f for f in received_down_frames if isinstance(f, TranscriptionFrame)
|
||||||
|
]
|
||||||
|
self.assertEqual(len(transcription_frames), 1)
|
||||||
|
self.assertEqual(transcription_frames[0].text, "DTMF: 0123456789*#")
|
||||||
Reference in New Issue
Block a user