DtmfAggregator: cancel interruption task to avoid a dangling task
This commit is contained in:
@@ -64,6 +64,7 @@ class DTMFAggregator(FrameProcessor):
|
|||||||
|
|
||||||
self._digit_event = asyncio.Event()
|
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 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.
|
||||||
@@ -81,6 +82,7 @@ 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
|
||||||
@@ -100,7 +102,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:
|
||||||
asyncio.create_task(self._send_interruption_task())
|
self._interruption_task = self.create_task(self._send_interruption_task())
|
||||||
|
|
||||||
# Check for immediate flush conditions
|
# Check for immediate flush conditions
|
||||||
if frame.button == self._termination_digit:
|
if frame.button == self._termination_digit:
|
||||||
@@ -111,12 +113,13 @@ class DTMFAggregator(FrameProcessor):
|
|||||||
|
|
||||||
async def _send_interruption_task(self):
|
async def _send_interruption_task(self):
|
||||||
"""Send interruption frame safely in a separate task."""
|
"""Send interruption frame safely in a separate task."""
|
||||||
try:
|
await self.push_frame(BotInterruptionFrame(), FrameDirection.UPSTREAM)
|
||||||
# Send the interruption frame
|
|
||||||
await self.push_frame(BotInterruptionFrame(), FrameDirection.UPSTREAM)
|
async def _stop_interruption_task(self) -> None:
|
||||||
except Exception as e:
|
"""Stops the interruption task."""
|
||||||
# Log error but don't propagate
|
if self._interruption_task:
|
||||||
print(f"Error sending interruption: {e}")
|
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."""
|
||||||
|
|||||||
@@ -214,7 +214,7 @@ class TestDTMFAggregator(unittest.IsolatedAsyncioTestCase):
|
|||||||
]
|
]
|
||||||
|
|
||||||
# All the InputDTMFFrames plus one TranscriptionFrame
|
# All the InputDTMFFrames plus one TranscriptionFrame
|
||||||
expected_down_frames = [InputDTMFFrame] * 12 + [TranscriptionFrame]
|
expected_down_frames = [InputDTMFFrame] * len(frames_to_send) + [TranscriptionFrame]
|
||||||
|
|
||||||
received_down_frames, _ = await run_test(
|
received_down_frames, _ = await run_test(
|
||||||
aggregator,
|
aggregator,
|
||||||
|
|||||||
Reference in New Issue
Block a user