Improve user turn stop timing by triggering timeout from VAD stop
Refactor TranscriptionUserTurnStopStrategy and TurnAnalyzerUserTurnStopStrategy to use VADUserStoppedSpeakingFrame as the ground truth for when speech ended, rather than triggering timeouts from transcription frames.
This commit is contained in:
@@ -12,6 +12,9 @@ from dataclasses import dataclass
|
||||
from pipecat.frames.frames import (
|
||||
Frame,
|
||||
ManuallySwitchServiceFrame,
|
||||
RequestMetadataFrame,
|
||||
ServiceMetadataFrame,
|
||||
StartFrame,
|
||||
SystemFrame,
|
||||
TextFrame,
|
||||
)
|
||||
@@ -54,6 +57,47 @@ class MockFrameProcessor(FrameProcessor):
|
||||
self.frame_count = 0
|
||||
|
||||
|
||||
@dataclass
|
||||
class MockMetadataFrame(ServiceMetadataFrame):
|
||||
"""A mock metadata frame for testing ServiceMetadataFrame handling."""
|
||||
|
||||
pass
|
||||
|
||||
|
||||
class MockMetadataService(FrameProcessor):
|
||||
"""A mock service that emits ServiceMetadataFrame like STT services.
|
||||
|
||||
Pushes MockMetadataFrame on StartFrame and RequestMetadataFrame.
|
||||
"""
|
||||
|
||||
def __init__(self, test_name: str, **kwargs):
|
||||
super().__init__(name=test_name, **kwargs)
|
||||
self.test_name = test_name
|
||||
self.processed_frames = []
|
||||
self.metadata_push_count = 0
|
||||
|
||||
async def process_frame(self, frame: Frame, direction: FrameDirection):
|
||||
await super().process_frame(frame, direction)
|
||||
self.processed_frames.append(frame)
|
||||
|
||||
if isinstance(frame, StartFrame):
|
||||
await self.push_frame(frame, direction)
|
||||
await self._push_metadata()
|
||||
elif isinstance(frame, RequestMetadataFrame):
|
||||
# Don't push RequestMetadataFrame downstream (it's internal)
|
||||
await self._push_metadata()
|
||||
else:
|
||||
await self.push_frame(frame, direction)
|
||||
|
||||
async def _push_metadata(self):
|
||||
self.metadata_push_count += 1
|
||||
await self.push_frame(MockMetadataFrame(service_name=self.test_name))
|
||||
|
||||
def reset_counters(self):
|
||||
self.processed_frames = []
|
||||
self.metadata_push_count = 0
|
||||
|
||||
|
||||
@dataclass
|
||||
class DummySystemFrame(SystemFrame):
|
||||
"""A dummy system frame for testing purposes."""
|
||||
@@ -336,5 +380,84 @@ class TestServiceSwitcher(unittest.IsolatedAsyncioTestCase):
|
||||
self.assertEqual(switcher2_service2_texts[0].text, "After switching second switcher")
|
||||
|
||||
|
||||
class TestServiceSwitcherMetadata(unittest.IsolatedAsyncioTestCase):
|
||||
"""Test cases for ServiceMetadataFrame handling in ServiceSwitcher."""
|
||||
|
||||
def setUp(self):
|
||||
"""Set up test fixtures with mock metadata services."""
|
||||
self.service1 = MockMetadataService("service1")
|
||||
self.service2 = MockMetadataService("service2")
|
||||
self.services = [self.service1, self.service2]
|
||||
|
||||
async def test_only_active_service_metadata_at_startup(self):
|
||||
"""Test that only the active service's metadata leaves the ServiceSwitcher at startup."""
|
||||
switcher = ServiceSwitcher(self.services, ServiceSwitcherStrategyManual)
|
||||
|
||||
# Run the pipeline (StartFrame triggers metadata emission)
|
||||
output_frames = []
|
||||
|
||||
async def capture_frame(frame: Frame):
|
||||
output_frames.append(frame)
|
||||
|
||||
await run_test(
|
||||
switcher,
|
||||
frames_to_send=[TextFrame(text="test")],
|
||||
expected_down_frames=[MockMetadataFrame, TextFrame],
|
||||
expected_up_frames=[],
|
||||
)
|
||||
|
||||
# Both services push metadata internally on StartFrame, but only the
|
||||
# active service's metadata passes through the filter
|
||||
self.assertEqual(self.service1.metadata_push_count, 1) # StartFrame (passes filter)
|
||||
self.assertEqual(self.service2.metadata_push_count, 1) # StartFrame (blocked by filter)
|
||||
|
||||
async def test_metadata_emitted_on_service_switch(self):
|
||||
"""Test that switching services triggers metadata emission from the new active service."""
|
||||
switcher = ServiceSwitcher(self.services, ServiceSwitcherStrategyManual)
|
||||
|
||||
# Reset counters after startup
|
||||
self.service1.reset_counters()
|
||||
self.service2.reset_counters()
|
||||
|
||||
await run_test(
|
||||
switcher,
|
||||
frames_to_send=[
|
||||
TextFrame(text="before switch"),
|
||||
ManuallySwitchServiceFrame(service=self.service2),
|
||||
TextFrame(text="after switch"),
|
||||
],
|
||||
expected_down_frames=[
|
||||
MockMetadataFrame, # From startup (service1)
|
||||
TextFrame,
|
||||
ManuallySwitchServiceFrame,
|
||||
TextFrame,
|
||||
MockMetadataFrame, # From service2 after switch
|
||||
],
|
||||
expected_up_frames=[],
|
||||
)
|
||||
|
||||
# service2 should have received RequestMetadataFrame after becoming active
|
||||
request_frames = [
|
||||
f for f in self.service2.processed_frames if isinstance(f, RequestMetadataFrame)
|
||||
]
|
||||
self.assertEqual(len(request_frames), 1)
|
||||
|
||||
async def test_inactive_service_metadata_blocked(self):
|
||||
"""Test that metadata from inactive services is blocked."""
|
||||
switcher = ServiceSwitcher(self.services, ServiceSwitcherStrategyManual)
|
||||
|
||||
# Run and collect output frames
|
||||
await run_test(
|
||||
switcher,
|
||||
frames_to_send=[TextFrame(text="test")],
|
||||
expected_down_frames=[MockMetadataFrame, TextFrame],
|
||||
expected_up_frames=[],
|
||||
)
|
||||
|
||||
# service2 pushed metadata on StartFrame, but it should have been blocked
|
||||
self.assertGreaterEqual(self.service2.metadata_push_count, 1)
|
||||
# Only one MockMetadataFrame should have left (from service1)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
|
||||
Reference in New Issue
Block a user