LLMUserAggregator: add bot turn start strategies timeout fallback

This commit is contained in:
Aleix Conchillo Flaqué
2025-12-23 13:56:21 -08:00
parent 40493e8ce8
commit 1f0357ae5e
2 changed files with 74 additions and 9 deletions

4
changelog/3291.added.md Normal file
View File

@@ -0,0 +1,4 @@
- `LLMUserAggregator` now exposes the following events:
- `on_user_turn_started`: triggered when a user turn starts
- `on_bot_turn_started`: triggered when a user turn ends and a bot turn starts
- `on_user_turn_end_timeout`: triggered when a user turn does not stop and times out

View File

@@ -51,6 +51,8 @@ from pipecat.frames.frames import (
UserImageRawFrame, UserImageRawFrame,
UserStartedSpeakingFrame, UserStartedSpeakingFrame,
UserStoppedSpeakingFrame, UserStoppedSpeakingFrame,
VADUserStartedSpeakingFrame,
VADUserStoppedSpeakingFrame,
) )
from pipecat.processors.aggregators.llm_context import ( from pipecat.processors.aggregators.llm_context import (
LLMContext, LLMContext,
@@ -75,9 +77,12 @@ class LLMUserAggregatorParams:
interruption frames. This is enabled by default, but you may want interruption frames. This is enabled by default, but you may want
to disable it if another component (e.g., an STT service) is already to disable it if another component (e.g., an STT service) is already
generating these frames. generating these frames.
user_turn_end_timeout: Time in seconds to wait before considering the
user's turn finished and starting the bot turn.
""" """
enable_user_speaking_frames: bool = True enable_user_speaking_frames: bool = True
user_turn_end_timeout: float = 5.0
@dataclass @dataclass
@@ -226,15 +231,20 @@ class LLMUserAggregator(LLMContextAggregator):
- on_user_turn_started: Called when the user turn starts - on_user_turn_started: Called when the user turn starts
- on_bot_turn_started: Called when the user turn ends and it is now the bots turn - on_bot_turn_started: Called when the user turn ends and it is now the bots turn
- on_user_turn_end_timeout: Called when no bot turn start strategy triggers
Example:: Example::
@aggregator.event_handler("on_user_turn_started") @aggregator.event_handler("on_user_turn_started")
async def on_user_turn_started(aggregator, strategy): async def on_user_turn_started(aggregator, Optional[strategy]):
... ...
@aggregator.event_handler("on_bot_turn_started") @aggregator.event_handler("on_bot_turn_started")
async def on_bot_turn_started(aggregator, strategy): async def on_bot_turn_started(aggregator, Optional[strategy]):
...
@aggregator.event_handler("on_user_turn_end_timeout")
async def on_user_turn_end_timeout(aggregator):
... ...
""" """
@@ -255,9 +265,15 @@ class LLMUserAggregator(LLMContextAggregator):
""" """
super().__init__(context=context, role="user", **kwargs) super().__init__(context=context, role="user", **kwargs)
self._params = params or LLMUserAggregatorParams() self._params = params or LLMUserAggregatorParams()
self._user_speaking = False
self._vad_user_speaking = False
self._user_turn = False
self._user_turn_end_timeout_event = asyncio.Event()
self._user_turn_end_timeout_task: Optional[asyncio.Task] = None
self._register_event_handler("on_user_turn_started") self._register_event_handler("on_user_turn_started")
self._register_event_handler("on_user_turn_end_timeout")
self._register_event_handler("on_bot_turn_started") self._register_event_handler("on_bot_turn_started")
async def cleanup(self): async def cleanup(self):
@@ -299,6 +315,12 @@ class LLMUserAggregator(LLMContextAggregator):
elif isinstance(frame, CancelFrame): elif isinstance(frame, CancelFrame):
await self._cancel(frame) await self._cancel(frame)
await self.push_frame(frame, direction) await self.push_frame(frame, direction)
elif isinstance(frame, VADUserStartedSpeakingFrame):
await self._handle_vad_user_started_speaking(frame)
await self.push_frame(frame, direction)
elif isinstance(frame, VADUserStoppedSpeakingFrame):
await self._handle_vad_user_stopped_speaking(frame)
await self.push_frame(frame, direction)
elif isinstance(frame, TranscriptionFrame): elif isinstance(frame, TranscriptionFrame):
await self._handle_transcription(frame) await self._handle_transcription(frame)
elif isinstance(frame, LLMRunFrame): elif isinstance(frame, LLMRunFrame):
@@ -335,6 +357,11 @@ class LLMUserAggregator(LLMContextAggregator):
await self.push_context_frame() await self.push_context_frame()
async def _start(self, frame: StartFrame): async def _start(self, frame: StartFrame):
if not self._user_turn_end_timeout_task:
self._user_turn_end_timeout_task = self.create_task(
self._user_turn_end_timeout_task_handler()
)
if self.turn_start_strategies and self.turn_start_strategies.user: if self.turn_start_strategies and self.turn_start_strategies.user:
for s in self.turn_start_strategies.user: for s in self.turn_start_strategies.user:
await s.setup(self.task_manager) await s.setup(self.task_manager)
@@ -356,6 +383,10 @@ class LLMUserAggregator(LLMContextAggregator):
await self._cleanup() await self._cleanup()
async def _cleanup(self): async def _cleanup(self):
if self._user_turn_end_timeout_task:
await self.cancel_task(self._user_turn_end_timeout_task)
self._user_turn_end_timeout_task = None
if self.turn_start_strategies and self.turn_start_strategies.user: if self.turn_start_strategies and self.turn_start_strategies.user:
for s in self.turn_start_strategies.user: for s in self.turn_start_strategies.user:
await s.cleanup() await s.cleanup()
@@ -406,6 +437,18 @@ class LLMUserAggregator(LLMContextAggregator):
" )" " )"
) )
async def _handle_vad_user_started_speaking(self, frame: VADUserStartedSpeakingFrame):
self._vad_user_speaking = True
# The user started talking, let's reset the user turn timeout.
self._user_turn_end_timeout_event.set()
async def _handle_vad_user_stopped_speaking(self, frame: VADUserStoppedSpeakingFrame):
self._vad_user_speaking = False
# The user stopped talking, let's reset the user turn timeout.
self._user_turn_end_timeout_event.set()
async def _handle_transcription(self, frame: TranscriptionFrame): async def _handle_transcription(self, frame: TranscriptionFrame):
text = frame.text text = frame.text
@@ -413,6 +456,9 @@ class LLMUserAggregator(LLMContextAggregator):
if not text.strip(): if not text.strip():
return return
# We have creceived a transcription, let's reset the user turn timeout.
self._user_turn_end_timeout_event.set()
# Transcriptions never include inter-part spaces (so far). # Transcriptions never include inter-part spaces (so far).
self._aggregation.append( self._aggregation.append(
TextPartForConcatenation( TextPartForConcatenation(
@@ -442,12 +488,13 @@ class LLMUserAggregator(LLMContextAggregator):
): ):
await self.broadcast_frame(frame_cls, **kwargs) await self.broadcast_frame(frame_cls, **kwargs)
async def _trigger_user_turn_start(self, strategy: BaseUserTurnStartStrategy): async def _trigger_user_turn_start(self, strategy: Optional[BaseUserTurnStartStrategy]):
# Prevent two consecutive user turn starts. # Prevent two consecutive user turn starts.
if self._user_speaking: if self._user_turn:
return return
self._user_speaking = True self._user_turn = True
self._user_turn_end_timeout_event.set()
# Reset all user turn start strategies to start fresh. # Reset all user turn start strategies to start fresh.
if self.turn_start_strategies and self.turn_start_strategies.user: if self.turn_start_strategies and self.turn_start_strategies.user:
@@ -462,12 +509,13 @@ class LLMUserAggregator(LLMContextAggregator):
await self._call_event_handler("on_user_turn_started", strategy) await self._call_event_handler("on_user_turn_started", strategy)
async def _trigger_bot_turn_start(self, strategy: BaseBotTurnStartStrategy): async def _trigger_bot_turn_start(self, strategy: Optional[BaseBotTurnStartStrategy]):
# Prevent two consecutive bot turn starts. # Prevent two consecutive bot turn starts.
if not self._user_speaking: if not self._user_turn:
return return
self._user_speaking = False self._user_turn = False
self._user_turn_end_timeout_event.set()
# Reset all bot turn start strategies to start fresh. # Reset all bot turn start strategies to start fresh.
if self.turn_start_strategies and self.turn_start_strategies.bot: if self.turn_start_strategies and self.turn_start_strategies.bot:
@@ -484,6 +532,19 @@ class LLMUserAggregator(LLMContextAggregator):
# Always push context frame. # Always push context frame.
await self.push_aggregation() await self.push_aggregation()
async def _user_turn_end_timeout_task_handler(self):
while True:
try:
await asyncio.wait_for(
self._user_turn_end_timeout_event.wait(),
timeout=self._params.user_turn_end_timeout,
)
self._user_turn_end_timeout_event.clear()
except asyncio.TimeoutError:
if self._user_turn and not self._vad_user_speaking:
await self._call_event_handler("on_user_turn_end_timeout")
await self._trigger_bot_turn_start(None)
class LLMAssistantAggregator(LLMContextAggregator): class LLMAssistantAggregator(LLMContextAggregator):
"""Assistant LLM aggregator that processes bot responses and function calls. """Assistant LLM aggregator that processes bot responses and function calls.