fix: cancel pipeline on peer disconnect
This commit is contained in:
@@ -7,9 +7,13 @@ from unittest.mock import AsyncMock, patch
|
||||
from pipecat.frames.frames import (
|
||||
BotStartedSpeakingFrame,
|
||||
BotStoppedSpeakingFrame,
|
||||
CancelFrame,
|
||||
OutputTransportMessageUrgentFrame,
|
||||
)
|
||||
from services.pipecat.pipeline_events import bind_cascade_pipeline_events
|
||||
from services.pipecat.pipeline_events import (
|
||||
bind_cascade_pipeline_events,
|
||||
bind_realtime_pipeline_events,
|
||||
)
|
||||
|
||||
|
||||
class _EventSource:
|
||||
@@ -84,6 +88,49 @@ class _Brain:
|
||||
|
||||
|
||||
class PipelineEventTest(unittest.IsolatedAsyncioTestCase):
|
||||
async def test_cascade_disconnect_cancels_instead_of_draining_audio(self):
|
||||
transport = _EventSource()
|
||||
worker = _Worker()
|
||||
brain = _Brain(worker)
|
||||
|
||||
bind_cascade_pipeline_events(
|
||||
transport=transport,
|
||||
worker=worker,
|
||||
brain=brain,
|
||||
context=SimpleNamespace(),
|
||||
text_input=_EventSource(),
|
||||
user_aggregator=_EventSource(),
|
||||
assistant_aggregator=_EventSource(),
|
||||
greeting="",
|
||||
vision_enabled=False,
|
||||
vision_state={"client_id": None},
|
||||
)
|
||||
|
||||
await transport.handlers["on_client_disconnected"](transport, object())
|
||||
|
||||
self.assertIsInstance(worker.frames[-1], CancelFrame)
|
||||
self.assertEqual(worker.frames[-1].reason, "client_disconnected")
|
||||
|
||||
async def test_realtime_disconnect_cancels_instead_of_draining_audio(self):
|
||||
transport = _EventSource()
|
||||
worker = _Worker()
|
||||
|
||||
bind_realtime_pipeline_events(
|
||||
transport=transport,
|
||||
worker=worker,
|
||||
realtime=SimpleNamespace(),
|
||||
brain=SimpleNamespace(),
|
||||
text_input=_EventSource(),
|
||||
greeting="",
|
||||
vision_enabled=False,
|
||||
vision_state={"client_id": None},
|
||||
)
|
||||
|
||||
await transport.handlers["on_client_disconnected"](transport, object())
|
||||
|
||||
self.assertIsInstance(worker.frames[-1], CancelFrame)
|
||||
self.assertEqual(worker.frames[-1].reason, "client_disconnected")
|
||||
|
||||
async def test_interruption_acknowledges_deferred_client_tool_result(self):
|
||||
transport = _EventSource()
|
||||
text_input = _EventSource()
|
||||
|
||||
Reference in New Issue
Block a user