LLMUserAggregator: add bot turn start strategies timeout fallback
This commit is contained in:
4
changelog/3291.added.md
Normal file
4
changelog/3291.added.md
Normal 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
|
||||||
@@ -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 bot’s turn
|
- on_bot_turn_started: Called when the user turn ends and it is now the bot’s 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.
|
||||||
|
|||||||
Reference in New Issue
Block a user