fix: interrupt message output before tool result

This commit is contained in:
Xin Wang
2026-08-04 09:30:52 +08:00
parent d068927b53
commit a16ecd8e01
10 changed files with 152 additions and 37 deletions

View File

@@ -1541,8 +1541,8 @@ class WorkflowBrainTests(unittest.IsolatedAsyncioTestCase):
):
events.append("message_completed")
async def interrupt_output():
events.append("interrupted")
async def wait_for_output_stopped():
events.append("output_stopped")
class OrderedCallEnd(FakeCallEnd):
def __init__(self):
@@ -1580,7 +1580,7 @@ class WorkflowBrainTests(unittest.IsolatedAsyncioTestCase):
call_end=call_end,
client_tools=client_tools,
set_input_enabled=input_states.append,
interrupt_output=interrupt_output,
wait_for_output_stopped=wait_for_output_stopped,
)
brain._message_stages.set_client_tools(client_tools)
@@ -1598,6 +1598,7 @@ class WorkflowBrainTests(unittest.IsolatedAsyncioTestCase):
self.assertEqual(input_states, [False])
self.assertEqual(client_tools.function_name, "show_message")
self.assertEqual(client_tools.options["response_wait_mode"], "session")
self.assertTrue(client_tools.options["interrupt_on_result"])
user_confirmed.set()
result = await message_task
@@ -1605,7 +1606,7 @@ class WorkflowBrainTests(unittest.IsolatedAsyncioTestCase):
self.assertEqual(result.action, "confirmed")
self.assertFalse(call_end.playback_completion.done())
self.assertEqual(input_states, [False, True])
self.assertEqual(events[-2:], ["interrupted", "message_completed"])
self.assertEqual(events[-2:], ["output_stopped", "message_completed"])
async def test_speech_only_message_waits_for_transport_playback(self):
brain = WorkflowBrain(

View File

@@ -1,5 +1,6 @@
from __future__ import annotations
import asyncio
import unittest
from pipecat.frames.frames import BotStartedSpeakingFrame, BotStoppedSpeakingFrame
@@ -36,6 +37,20 @@ class CallEndCoordinatorTest(unittest.IsolatedAsyncioTestCase):
self.assertEqual(self.reasons, ["prompt_end_call"])
async def test_wait_until_silent_tracks_transport_boundary(self):
self.assertFalse(self.coordinator.speaking)
await self.coordinator.wait_until_silent()
await self.coordinator.observe(BotStartedSpeakingFrame())
self.assertTrue(self.coordinator.speaking)
wait_task = asyncio.create_task(self.coordinator.wait_until_silent())
await asyncio.sleep(0)
self.assertFalse(wait_task.done())
await self.coordinator.observe(BotStoppedSpeakingFrame())
await wait_task
self.assertFalse(self.coordinator.speaking)
async def test_tool_only_end_call_finishes_without_waiting(self):
self.coordinator.begin_response()
self.coordinator.begin("tool_only")

View File

@@ -138,6 +138,51 @@ class ClientToolExecutorTests(unittest.IsolatedAsyncioTestCase):
class ClientToolBrokerTests(unittest.IsolatedAsyncioTestCase):
async def test_interrupts_before_resolving_configured_result(self):
broker = ClientToolBroker()
outbound = []
observed_future_states = []
async def push_frame(frame, direction=FrameDirection.DOWNSTREAM):
outbound.append((frame, direction))
async def broadcast_interruption():
observed_future_states.append(
[pending.future.done() for pending in broker._pending.values()]
)
broker.push_frame = push_frame
broker.broadcast_interruption = broadcast_interruption
call = asyncio.create_task(
broker.call(
"show_message",
{},
timeout_seconds=1,
response_wait_mode="session",
interrupt_on_result=True,
)
)
await asyncio.sleep(0)
message = outbound[0][0].message
await broker.process_frame(
InputTransportMessageFrame(
message={
"type": "client-tool-result",
"tool_call_id": message["tool_call_id"],
"status": "ok",
"data": {"action": "confirmed"},
}
),
FrameDirection.DOWNSTREAM,
)
self.assertEqual(observed_future_states, [[False]])
self.assertEqual(
await call,
{"status": "ok", "data": {"action": "confirmed"}},
)
async def test_correlates_result_and_times_out(self):
broker = ClientToolBroker()
outbound = []

View File

@@ -10,7 +10,7 @@ from pipecat.frames.frames import (
OutputTransportMessageUrgentFrame,
)
from services.pipecat.pipeline_events import bind_cascade_pipeline_events
from services.pipecat.pipeline import _interrupt_pipeline_output
from services.pipecat.pipeline import _wait_for_interrupted_output
class _EventSource:
@@ -81,24 +81,31 @@ class _Brain:
class PipelineEventTest(unittest.IsolatedAsyncioTestCase):
async def test_output_interruption_broadcasts_before_flush_barrier(self):
async def test_interrupted_output_waits_for_flush_and_stop(self):
events = []
source = SimpleNamespace(
broadcast_interruption=AsyncMock(
side_effect=lambda: events.append("broadcast")
)
)
async def wait_until_stopped():
events.append("stopped")
worker = SimpleNamespace(
flush_pipeline=AsyncMock(
side_effect=lambda **_kwargs: events.append("flush")
side_effect=lambda **_kwargs: events.append("flush") or True
)
)
await _interrupt_pipeline_output(source, worker)
await _wait_for_interrupted_output(
worker,
wait_until_stopped=wait_until_stopped,
)
self.assertEqual(events, ["broadcast", "flush"])
source.broadcast_interruption.assert_awaited_once_with()
worker.flush_pipeline.assert_awaited_once_with(timeout=1.0)
self.assertEqual(events, ["flush", "stopped"])
worker.flush_pipeline.assert_awaited_once_with(timeout=2.0)
async def test_interrupted_output_rejects_flush_timeout(self):
worker = SimpleNamespace(flush_pipeline=AsyncMock(return_value=False))
with self.assertRaisesRegex(RuntimeError, "中断帧"):
await _wait_for_interrupted_output(worker)
async def test_greeting_keeps_playback_timestamp_until_client_ready(self):
transport = _EventSource()