Merge pull request #2479 from pipecat-ai/aleix/simplify-dtmf-aggregator
DTMFAggregator: no need for interruption task
This commit is contained in:
@@ -24,8 +24,9 @@ from pipecat.frames.frames import (
|
|||||||
StartFrame,
|
StartFrame,
|
||||||
TranscriptionFrame,
|
TranscriptionFrame,
|
||||||
)
|
)
|
||||||
from pipecat.processors.frame_processor import FrameDirection, FrameProcessor
|
from pipecat.processors.frame_processor import FrameDirection, FrameProcessor, FrameProcessorSetup
|
||||||
from pipecat.utils.asyncio.timeout import wait_for
|
from pipecat.utils.asyncio.timeout import wait_for
|
||||||
|
from pipecat.utils.asyncio.watchdog_event import WatchdogEvent
|
||||||
from pipecat.utils.time import time_now_iso8601
|
from pipecat.utils.time import time_now_iso8601
|
||||||
|
|
||||||
|
|
||||||
@@ -63,9 +64,17 @@ class DTMFAggregator(FrameProcessor):
|
|||||||
self._termination_digit = termination_digit
|
self._termination_digit = termination_digit
|
||||||
self._prefix = prefix
|
self._prefix = prefix
|
||||||
|
|
||||||
self._digit_event = asyncio.Event()
|
|
||||||
self._aggregation_task: Optional[asyncio.Task] = None
|
self._aggregation_task: Optional[asyncio.Task] = None
|
||||||
self._interruption_task: Optional[asyncio.Task] = None
|
|
||||||
|
async def setup(self, setup: FrameProcessorSetup):
|
||||||
|
"""Setup resources."""
|
||||||
|
await super().setup(setup)
|
||||||
|
self._digit_event = WatchdogEvent(setup.task_manager)
|
||||||
|
|
||||||
|
async def cleanup(self) -> None:
|
||||||
|
"""Clean up resources."""
|
||||||
|
await super().cleanup()
|
||||||
|
await self._stop_aggregation_task()
|
||||||
|
|
||||||
async def process_frame(self, frame: Frame, direction: FrameDirection) -> None:
|
async def process_frame(self, frame: Frame, direction: FrameDirection) -> None:
|
||||||
"""Process incoming frames and handle DTMF aggregation.
|
"""Process incoming frames and handle DTMF aggregation.
|
||||||
@@ -83,7 +92,6 @@ class DTMFAggregator(FrameProcessor):
|
|||||||
if self._aggregation:
|
if self._aggregation:
|
||||||
await self._flush_aggregation()
|
await self._flush_aggregation()
|
||||||
await self._stop_aggregation_task()
|
await self._stop_aggregation_task()
|
||||||
await self._stop_interruption_task()
|
|
||||||
await self.push_frame(frame, direction)
|
await self.push_frame(frame, direction)
|
||||||
elif isinstance(frame, InputDTMFFrame):
|
elif isinstance(frame, InputDTMFFrame):
|
||||||
# Push the DTMF frame downstream first
|
# Push the DTMF frame downstream first
|
||||||
@@ -103,7 +111,7 @@ class DTMFAggregator(FrameProcessor):
|
|||||||
|
|
||||||
# For first digit, schedule interruption in separate task
|
# For first digit, schedule interruption in separate task
|
||||||
if is_first_digit:
|
if is_first_digit:
|
||||||
self._interruption_task = self.create_task(self._send_interruption_task())
|
await self.push_frame(BotInterruptionFrame(), FrameDirection.UPSTREAM)
|
||||||
|
|
||||||
# Check for immediate flush conditions
|
# Check for immediate flush conditions
|
||||||
if frame.button == self._termination_digit:
|
if frame.button == self._termination_digit:
|
||||||
@@ -112,16 +120,6 @@ class DTMFAggregator(FrameProcessor):
|
|||||||
# Signal digit received for timeout handling
|
# Signal digit received for timeout handling
|
||||||
self._digit_event.set()
|
self._digit_event.set()
|
||||||
|
|
||||||
async def _send_interruption_task(self):
|
|
||||||
"""Send interruption frame safely in a separate task."""
|
|
||||||
await self.push_frame(BotInterruptionFrame(), FrameDirection.UPSTREAM)
|
|
||||||
|
|
||||||
async def _stop_interruption_task(self) -> None:
|
|
||||||
"""Stops the interruption task."""
|
|
||||||
if self._interruption_task:
|
|
||||||
await self.cancel_task(self._interruption_task)
|
|
||||||
self._interruption_task = None
|
|
||||||
|
|
||||||
def _create_aggregation_task(self) -> None:
|
def _create_aggregation_task(self) -> None:
|
||||||
"""Creates the aggregation task if it hasn't been created yet."""
|
"""Creates the aggregation task if it hasn't been created yet."""
|
||||||
if not self._aggregation_task:
|
if not self._aggregation_task:
|
||||||
@@ -158,8 +156,3 @@ class DTMFAggregator(FrameProcessor):
|
|||||||
await self.push_frame(transcription_frame)
|
await self.push_frame(transcription_frame)
|
||||||
|
|
||||||
self._aggregation = ""
|
self._aggregation = ""
|
||||||
|
|
||||||
async def cleanup(self) -> None:
|
|
||||||
"""Clean up resources."""
|
|
||||||
await super().cleanup()
|
|
||||||
await self._stop_aggregation_task()
|
|
||||||
|
|||||||
@@ -29,12 +29,13 @@ class BaseObject(ABC):
|
|||||||
classes in the framework should inherit from this base class.
|
classes in the framework should inherit from this base class.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
def __init__(self, *, name: Optional[str] = None):
|
def __init__(self, *, name: Optional[str] = None, **kwargs):
|
||||||
"""Initialize the base object.
|
"""Initialize the base object.
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
name: Optional custom name for the object. If not provided,
|
name: Optional custom name for the object. If not provided,
|
||||||
generates a name using the class name and instance count.
|
generates a name using the class name and instance count.
|
||||||
|
**kwargs: Additional arguments passed to parent class.
|
||||||
"""
|
"""
|
||||||
self._id: int = obj_id()
|
self._id: int = obj_id()
|
||||||
self._name = name or f"{self.__class__.__name__}#{obj_count(self)}"
|
self._name = name or f"{self.__class__.__name__}#{obj_count(self)}"
|
||||||
|
|||||||
Reference in New Issue
Block a user