diff --git a/backend/services/pipecat/pipeline_events.py b/backend/services/pipecat/pipeline_events.py index ef8c782..c414a17 100644 --- a/backend/services/pipecat/pipeline_events.py +++ b/backend/services/pipecat/pipeline_events.py @@ -7,7 +7,7 @@ from loguru import logger from pipecat.frames.frames import ( BotStartedSpeakingFrame, BotStoppedSpeakingFrame, - EndFrame, + CancelFrame, OutputTransportMessageUrgentFrame, TTSSpeakFrame, ) @@ -218,8 +218,11 @@ def bind_cascade_pipeline_events( @transport.event_handler("on_client_disconnected") async def on_client_disconnected(_transport, _client): - logger.info("对端断开,结束管线") - await worker.queue_frame(EndFrame()) + # The peer can no longer consume queued audio. EndFrame would wait for + # graceful playback and can deadlock on SmallWebRTC's pending audio + # future, so a transport disconnect must stop output immediately. + logger.info("对端断开,立即取消管线") + await worker.queue_frame(CancelFrame(reason="client_disconnected")) def bind_realtime_pipeline_events( @@ -296,5 +299,5 @@ def bind_realtime_pipeline_events( @transport.event_handler("on_client_disconnected") async def on_client_disconnected(_transport, _client): - logger.info("Realtime 对端断开,结束管线") - await worker.queue_frame(EndFrame()) + logger.info("Realtime 对端断开,立即取消管线") + await worker.queue_frame(CancelFrame(reason="client_disconnected")) diff --git a/backend/tests/test_pipeline_events.py b/backend/tests/test_pipeline_events.py index 2230d2f..dbaf47f 100644 --- a/backend/tests/test_pipeline_events.py +++ b/backend/tests/test_pipeline_events.py @@ -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()