fix: interrupt message output before tool result
This commit is contained in:
@@ -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(
|
||||
|
||||
@@ -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")
|
||||
|
||||
@@ -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 = []
|
||||
|
||||
@@ -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()
|
||||
|
||||
Reference in New Issue
Block a user