Merge pull request #3575 from pipecat-ai/mb/fix-turn-stopped-event-end-cancel-frame
Emit on_assistant_turn_stopped and on_user_turn_stopped from EndFrame…
This commit is contained in:
1
changelog/3575.fixed.md
Normal file
1
changelog/3575.fixed.md
Normal file
@@ -0,0 +1 @@
|
|||||||
|
- Fixed `LLMUserAggregator` and `LLMAssistantAggregator` not emitting pending transcripts via `on_user_turn_stopped` and `on_assistant_turn_stopped` events when the conversation ends (`EndFrame`) or is cancelled (`CancelFrame`).
|
||||||
@@ -464,9 +464,11 @@ class LLMUserAggregator(LLMContextAggregator):
|
|||||||
await s.setup(self.task_manager)
|
await s.setup(self.task_manager)
|
||||||
|
|
||||||
async def _stop(self, frame: EndFrame):
|
async def _stop(self, frame: EndFrame):
|
||||||
|
await self._maybe_emit_user_turn_stopped(on_session_end=True)
|
||||||
await self._cleanup()
|
await self._cleanup()
|
||||||
|
|
||||||
async def _cancel(self, frame: CancelFrame):
|
async def _cancel(self, frame: CancelFrame):
|
||||||
|
await self._maybe_emit_user_turn_stopped(on_session_end=True)
|
||||||
await self._cleanup()
|
await self._cleanup()
|
||||||
|
|
||||||
async def _cleanup(self):
|
async def _cleanup(self):
|
||||||
@@ -602,14 +604,7 @@ class LLMUserAggregator(LLMContextAggregator):
|
|||||||
if params.enable_user_speaking_frames:
|
if params.enable_user_speaking_frames:
|
||||||
await self.broadcast_frame(UserStoppedSpeakingFrame)
|
await self.broadcast_frame(UserStoppedSpeakingFrame)
|
||||||
|
|
||||||
# Always push context frame.
|
await self._maybe_emit_user_turn_stopped(strategy)
|
||||||
aggregation = await self.push_aggregation()
|
|
||||||
|
|
||||||
message = UserTurnStoppedMessage(
|
|
||||||
content=aggregation, timestamp=self._user_turn_start_timestamp
|
|
||||||
)
|
|
||||||
await self._call_event_handler("on_user_turn_stopped", strategy, message)
|
|
||||||
self._user_turn_start_timestamp = ""
|
|
||||||
|
|
||||||
async def _on_user_turn_stop_timeout(self, controller):
|
async def _on_user_turn_stop_timeout(self, controller):
|
||||||
await self._call_event_handler("on_user_turn_stop_timeout")
|
await self._call_event_handler("on_user_turn_stop_timeout")
|
||||||
@@ -617,6 +612,26 @@ class LLMUserAggregator(LLMContextAggregator):
|
|||||||
async def _on_user_turn_idle(self, controller):
|
async def _on_user_turn_idle(self, controller):
|
||||||
await self._call_event_handler("on_user_turn_idle")
|
await self._call_event_handler("on_user_turn_idle")
|
||||||
|
|
||||||
|
async def _maybe_emit_user_turn_stopped(
|
||||||
|
self,
|
||||||
|
strategy: Optional[BaseUserTurnStopStrategy] = None,
|
||||||
|
on_session_end: bool = False,
|
||||||
|
):
|
||||||
|
"""Maybe emit user turn stopped event.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
strategy: The strategy that triggered the turn stop.
|
||||||
|
on_session_end: If True, only emit if there's unemitted content
|
||||||
|
(avoids duplicate events when session ends).
|
||||||
|
"""
|
||||||
|
aggregation = await self.push_aggregation()
|
||||||
|
if not on_session_end or aggregation:
|
||||||
|
message = UserTurnStoppedMessage(
|
||||||
|
content=aggregation, timestamp=self._user_turn_start_timestamp
|
||||||
|
)
|
||||||
|
await self._call_event_handler("on_user_turn_stopped", strategy, message)
|
||||||
|
self._user_turn_start_timestamp = ""
|
||||||
|
|
||||||
|
|
||||||
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.
|
||||||
@@ -739,6 +754,9 @@ class LLMAssistantAggregator(LLMContextAggregator):
|
|||||||
if isinstance(frame, InterruptionFrame):
|
if isinstance(frame, InterruptionFrame):
|
||||||
await self._handle_interruptions(frame)
|
await self._handle_interruptions(frame)
|
||||||
await self.push_frame(frame, direction)
|
await self.push_frame(frame, direction)
|
||||||
|
elif isinstance(frame, (EndFrame, CancelFrame)):
|
||||||
|
await self._handle_end_or_cancel(frame)
|
||||||
|
await self.push_frame(frame, direction)
|
||||||
elif isinstance(frame, LLMFullResponseStartFrame):
|
elif isinstance(frame, LLMFullResponseStartFrame):
|
||||||
await self._handle_llm_start(frame)
|
await self._handle_llm_start(frame)
|
||||||
elif isinstance(frame, LLMFullResponseEndFrame):
|
elif isinstance(frame, LLMFullResponseEndFrame):
|
||||||
@@ -813,6 +831,10 @@ class LLMAssistantAggregator(LLMContextAggregator):
|
|||||||
self._started = 0
|
self._started = 0
|
||||||
await self.reset()
|
await self.reset()
|
||||||
|
|
||||||
|
async def _handle_end_or_cancel(self, frame: Frame):
|
||||||
|
await self._trigger_assistant_turn_stopped()
|
||||||
|
self._started = 0
|
||||||
|
|
||||||
async def _handle_function_calls_started(self, frame: FunctionCallsStartedFrame):
|
async def _handle_function_calls_started(self, frame: FunctionCallsStartedFrame):
|
||||||
function_names = [f"{f.function_name}:{f.tool_call_id}" for f in frame.function_calls]
|
function_names = [f"{f.function_name}:{f.tool_call_id}" for f in frame.function_calls]
|
||||||
logger.debug(f"{self} FunctionCallsStartedFrame: {function_names}")
|
logger.debug(f"{self} FunctionCallsStartedFrame: {function_names}")
|
||||||
|
|||||||
@@ -344,6 +344,35 @@ class TestLLMUserAggregator(unittest.IsolatedAsyncioTestCase):
|
|||||||
# The user mute strategies should have muted the user.
|
# The user mute strategies should have muted the user.
|
||||||
self.assertFalse(user_turn)
|
self.assertFalse(user_turn)
|
||||||
|
|
||||||
|
async def test_pending_transcription_emitted_on_end_frame(self):
|
||||||
|
"""Pending user transcription should be emitted when EndFrame arrives."""
|
||||||
|
context = LLMContext()
|
||||||
|
|
||||||
|
user_aggregator = LLMUserAggregator(context)
|
||||||
|
|
||||||
|
stop_messages = []
|
||||||
|
|
||||||
|
@user_aggregator.event_handler("on_user_turn_stopped")
|
||||||
|
async def on_user_turn_stopped(aggregator, strategy, message):
|
||||||
|
stop_messages.append((strategy, message))
|
||||||
|
|
||||||
|
pipeline = Pipeline([user_aggregator])
|
||||||
|
|
||||||
|
# Start turn and send transcription, but don't trigger normal turn stop
|
||||||
|
frames_to_send = [
|
||||||
|
VADUserStartedSpeakingFrame(),
|
||||||
|
TranscriptionFrame(text="Hello!", user_id="", timestamp="now"),
|
||||||
|
# No VADUserStoppedSpeakingFrame - turn doesn't stop normally
|
||||||
|
# EndFrame will be sent by run_test, triggering emission
|
||||||
|
]
|
||||||
|
await run_test(pipeline, frames_to_send=frames_to_send)
|
||||||
|
|
||||||
|
# The pending transcription should be emitted on EndFrame
|
||||||
|
self.assertEqual(len(stop_messages), 1)
|
||||||
|
strategy, message = stop_messages[0]
|
||||||
|
self.assertIsNone(strategy) # strategy is None for end/cancel
|
||||||
|
self.assertEqual(message.content, "Hello!")
|
||||||
|
|
||||||
|
|
||||||
class TestLLMAssistantAggregator(unittest.IsolatedAsyncioTestCase):
|
class TestLLMAssistantAggregator(unittest.IsolatedAsyncioTestCase):
|
||||||
async def test_empty(self):
|
async def test_empty(self):
|
||||||
@@ -512,3 +541,28 @@ class TestLLMAssistantAggregator(unittest.IsolatedAsyncioTestCase):
|
|||||||
]
|
]
|
||||||
await run_test(aggregator, frames_to_send=frames_to_send)
|
await run_test(aggregator, frames_to_send=frames_to_send)
|
||||||
self.assertEqual(thought_message.content, "I'm thinking!")
|
self.assertEqual(thought_message.content, "I'm thinking!")
|
||||||
|
|
||||||
|
async def test_pending_text_emitted_on_end_frame(self):
|
||||||
|
"""Pending assistant text should be emitted when EndFrame arrives."""
|
||||||
|
context = LLMContext()
|
||||||
|
|
||||||
|
aggregator = LLMAssistantAggregator(context)
|
||||||
|
|
||||||
|
stop_messages = []
|
||||||
|
|
||||||
|
@aggregator.event_handler("on_assistant_turn_stopped")
|
||||||
|
async def on_assistant_turn_stopped(aggregator, message: AssistantTurnStoppedMessage):
|
||||||
|
stop_messages.append(message)
|
||||||
|
|
||||||
|
# Start response and send text, but don't send LLMFullResponseEndFrame
|
||||||
|
frames_to_send = [
|
||||||
|
LLMFullResponseStartFrame(),
|
||||||
|
LLMTextFrame("Hello from Pipecat!"),
|
||||||
|
# No LLMFullResponseEndFrame - response doesn't end normally
|
||||||
|
# EndFrame will be sent by run_test, triggering emission
|
||||||
|
]
|
||||||
|
await run_test(aggregator, frames_to_send=frames_to_send)
|
||||||
|
|
||||||
|
# The pending text should be emitted on EndFrame
|
||||||
|
self.assertEqual(len(stop_messages), 1)
|
||||||
|
self.assertEqual(stop_messages[0].content, "Hello from Pipecat!")
|
||||||
|
|||||||
Reference in New Issue
Block a user