fix(workflow): track fixed speech at transport output
This commit is contained in:
@@ -3,14 +3,17 @@
|
|||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
from collections.abc import Awaitable
|
from collections.abc import Awaitable
|
||||||
from typing import Any
|
from typing import TYPE_CHECKING, Any
|
||||||
|
|
||||||
from pipecat.frames.frames import OutputTransportMessageUrgentFrame, TTSSpeakFrame
|
from pipecat.frames.frames import OutputTransportMessageUrgentFrame, TTSSpeakFrame
|
||||||
from pipecat.utils.time import time_now_iso8601
|
from pipecat.utils.time import time_now_iso8601
|
||||||
|
|
||||||
from services.brains.base import BrainRuntime
|
from services.pipecat.call_lifecycle import playback_marker_for
|
||||||
from services.runtime_variables import DynamicVariableStore
|
from services.runtime_variables import DynamicVariableStore
|
||||||
|
|
||||||
|
if TYPE_CHECKING:
|
||||||
|
from services.brains.base import BrainRuntime
|
||||||
|
|
||||||
|
|
||||||
FIXED_SPEECH_CONTEXT_MARKER = "[会话事实:以下固定消息已向用户播报]"
|
FIXED_SPEECH_CONTEXT_MARKER = "[会话事实:以下固定消息已向用户播报]"
|
||||||
|
|
||||||
@@ -79,6 +82,10 @@ class FixedSpeechOutput:
|
|||||||
await self._runtime.queue_frame(
|
await self._runtime.queue_frame(
|
||||||
TTSSpeakFrame(content, append_to_context=False)
|
TTSSpeakFrame(content, append_to_context=False)
|
||||||
)
|
)
|
||||||
|
playback_marker = playback_marker_for(playback_completion)
|
||||||
|
if playback_marker is not None:
|
||||||
|
await self._runtime.queue_frame(playback_marker)
|
||||||
|
playback_marker.completion.mark_queued()
|
||||||
return playback_completion
|
return playback_completion
|
||||||
|
|
||||||
async def emit(self, message: dict[str, Any]) -> None:
|
async def emit(self, message: dict[str, Any]) -> None:
|
||||||
|
|||||||
@@ -3,14 +3,66 @@
|
|||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
import asyncio
|
import asyncio
|
||||||
from collections import deque
|
|
||||||
from collections.abc import Awaitable, Callable
|
from collections.abc import Awaitable, Callable
|
||||||
|
from dataclasses import dataclass
|
||||||
|
|
||||||
from loguru import logger
|
from loguru import logger
|
||||||
from pipecat.frames.frames import BotStartedSpeakingFrame, BotStoppedSpeakingFrame
|
from pipecat.frames.frames import (
|
||||||
|
BotStartedSpeakingFrame,
|
||||||
|
BotStoppedSpeakingFrame,
|
||||||
|
DataFrame,
|
||||||
|
InterruptionFrame,
|
||||||
|
)
|
||||||
from pipecat.processors.frame_processor import FrameDirection, FrameProcessor
|
from pipecat.processors.frame_processor import FrameDirection, FrameProcessor
|
||||||
|
|
||||||
|
|
||||||
|
class SpeechPlaybackCompletion:
|
||||||
|
"""Awaitable completed by its exact marker at the transport output."""
|
||||||
|
|
||||||
|
def __init__(self, coordinator: CallEndCoordinator):
|
||||||
|
self._coordinator = coordinator
|
||||||
|
self._future: asyncio.Future[None] = (
|
||||||
|
asyncio.get_running_loop().create_future()
|
||||||
|
)
|
||||||
|
self._queued = False
|
||||||
|
|
||||||
|
def __await__(self):
|
||||||
|
return self._future.__await__()
|
||||||
|
|
||||||
|
def done(self) -> bool:
|
||||||
|
return self._future.done()
|
||||||
|
|
||||||
|
@property
|
||||||
|
def queued(self) -> bool:
|
||||||
|
return self._queued
|
||||||
|
|
||||||
|
def mark_queued(self) -> None:
|
||||||
|
self._queued = True
|
||||||
|
|
||||||
|
async def mark_played(self) -> None:
|
||||||
|
await self._coordinator.complete_tracked_speech(self)
|
||||||
|
|
||||||
|
def _resolve(self) -> None:
|
||||||
|
if not self._future.done():
|
||||||
|
self._future.set_result(None)
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class FixedSpeechPlaybackMarkerFrame(DataFrame):
|
||||||
|
"""Ordered frame that identifies one fixed utterance at transport output."""
|
||||||
|
|
||||||
|
completion: SpeechPlaybackCompletion
|
||||||
|
|
||||||
|
|
||||||
|
def playback_marker_for(
|
||||||
|
completion: Awaitable[None] | None,
|
||||||
|
) -> FixedSpeechPlaybackMarkerFrame | None:
|
||||||
|
"""Build a marker only for the production call-end coordinator."""
|
||||||
|
if not isinstance(completion, SpeechPlaybackCompletion):
|
||||||
|
return None
|
||||||
|
return FixedSpeechPlaybackMarkerFrame(completion=completion)
|
||||||
|
|
||||||
|
|
||||||
class CallEndCoordinator:
|
class CallEndCoordinator:
|
||||||
"""End immediately or after the currently armed closing speech finishes."""
|
"""End immediately or after the currently armed closing speech finishes."""
|
||||||
|
|
||||||
@@ -22,17 +74,10 @@ class CallEndCoordinator:
|
|||||||
self._speech_stopped = asyncio.Event()
|
self._speech_stopped = asyncio.Event()
|
||||||
self._speech_stopped.set()
|
self._speech_stopped.set()
|
||||||
self._response_speech_started = False
|
self._response_speech_started = False
|
||||||
self._tracked_speeches = 0
|
self._tracked_speech_completions: set[SpeechPlaybackCompletion] = set()
|
||||||
self._tracked_speech_completions: deque[asyncio.Future[None]] = deque()
|
|
||||||
self._finish_after_tracked_speech = False
|
self._finish_after_tracked_speech = False
|
||||||
self._finished = False
|
self._finished = False
|
||||||
self._reason = "completed"
|
self._reason = "completed"
|
||||||
# Only BotStoppedSpeakingFrame that follows a tracked BotStartedSpeakingFrame
|
|
||||||
# decrements the tracked-speech counter. This prevents a stale or cross-talk
|
|
||||||
# stop frame (e.g. from a preceding LLM utterance) from consuming the counter
|
|
||||||
# meant for a fixed end-node message.
|
|
||||||
self._pending_tracked_starts = 0
|
|
||||||
self._current_speech_is_tracked = False
|
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def ending(self) -> bool:
|
def ending(self) -> bool:
|
||||||
@@ -59,18 +104,32 @@ class CallEndCoordinator:
|
|||||||
"""Wait for the next observed bot speech to finish."""
|
"""Wait for the next observed bot speech to finish."""
|
||||||
self._armed = True
|
self._armed = True
|
||||||
|
|
||||||
def track_speech(self) -> Awaitable[None]:
|
def track_speech(self) -> SpeechPlaybackCompletion:
|
||||||
"""Register fixed speech and return its transport completion signal."""
|
"""Register fixed speech and return its transport completion signal."""
|
||||||
completion = asyncio.get_running_loop().create_future()
|
completion = SpeechPlaybackCompletion(self)
|
||||||
self._tracked_speech_completions.append(completion)
|
self._tracked_speech_completions.add(completion)
|
||||||
self._tracked_speeches += 1
|
|
||||||
self._pending_tracked_starts += 1
|
|
||||||
return completion
|
return completion
|
||||||
|
|
||||||
|
async def complete_tracked_speech(
|
||||||
|
self,
|
||||||
|
completion: SpeechPlaybackCompletion,
|
||||||
|
) -> None:
|
||||||
|
"""Complete one fixed utterance when its marker reaches output."""
|
||||||
|
if completion not in self._tracked_speech_completions:
|
||||||
|
return
|
||||||
|
self._tracked_speech_completions.remove(completion)
|
||||||
|
completion._resolve()
|
||||||
|
if (
|
||||||
|
self._finish_after_tracked_speech
|
||||||
|
and not self._tracked_speech_completions
|
||||||
|
):
|
||||||
|
logger.info("所有工作流结束语播报完毕,挂断通话")
|
||||||
|
await self.finish()
|
||||||
|
|
||||||
async def arm_after_tracked_speech(self) -> None:
|
async def arm_after_tracked_speech(self) -> None:
|
||||||
"""Finish after every already queued fixed utterance has played."""
|
"""Finish after every already queued fixed utterance has played."""
|
||||||
self._finish_after_tracked_speech = True
|
self._finish_after_tracked_speech = True
|
||||||
if self._tracked_speeches == 0:
|
if not self._tracked_speech_completions:
|
||||||
await self.finish()
|
await self.finish()
|
||||||
|
|
||||||
async def finish_after_current_speech(self, *, has_text: bool) -> None:
|
async def finish_after_current_speech(self, *, has_text: bool) -> None:
|
||||||
@@ -90,30 +149,25 @@ class CallEndCoordinator:
|
|||||||
await self._queue_end(self._reason)
|
await self._queue_end(self._reason)
|
||||||
|
|
||||||
async def observe(self, frame) -> None:
|
async def observe(self, frame) -> None:
|
||||||
if isinstance(frame, BotStartedSpeakingFrame):
|
if isinstance(frame, InterruptionFrame):
|
||||||
|
# Pipecat discards queued data frames on interruption, including
|
||||||
|
# playback markers. Treat already queued fixed speech as stopped so
|
||||||
|
# an interrupted Message cannot block a later EndNode forever.
|
||||||
|
interrupted = tuple(
|
||||||
|
completion
|
||||||
|
for completion in self._tracked_speech_completions
|
||||||
|
if completion.queued
|
||||||
|
)
|
||||||
|
for completion in interrupted:
|
||||||
|
await self.complete_tracked_speech(completion)
|
||||||
|
elif isinstance(frame, BotStartedSpeakingFrame):
|
||||||
self._speaking = True
|
self._speaking = True
|
||||||
self._speech_stopped.clear()
|
self._speech_stopped.clear()
|
||||||
self._response_speech_started = True
|
self._response_speech_started = True
|
||||||
if self._pending_tracked_starts > 0:
|
|
||||||
self._pending_tracked_starts -= 1
|
|
||||||
self._current_speech_is_tracked = True
|
|
||||||
elif isinstance(frame, BotStoppedSpeakingFrame) and self._speaking:
|
elif isinstance(frame, BotStoppedSpeakingFrame) and self._speaking:
|
||||||
self._speaking = False
|
self._speaking = False
|
||||||
self._speech_stopped.set()
|
self._speech_stopped.set()
|
||||||
if self._current_speech_is_tracked:
|
if self._armed:
|
||||||
self._current_speech_is_tracked = False
|
|
||||||
if self._tracked_speeches > 0:
|
|
||||||
self._tracked_speeches -= 1
|
|
||||||
completion = self._tracked_speech_completions.popleft()
|
|
||||||
if not completion.done():
|
|
||||||
completion.set_result(None)
|
|
||||||
if (
|
|
||||||
self._finish_after_tracked_speech
|
|
||||||
and self._tracked_speeches == 0
|
|
||||||
):
|
|
||||||
logger.info("所有工作流结束语播报完毕,挂断通话")
|
|
||||||
await self.finish()
|
|
||||||
elif self._armed:
|
|
||||||
logger.info("结束语播报完毕,挂断通话")
|
logger.info("结束语播报完毕,挂断通话")
|
||||||
await self.finish()
|
await self.finish()
|
||||||
|
|
||||||
|
|||||||
@@ -14,15 +14,67 @@ from pipecat.transports.base_transport import TransportParams
|
|||||||
|
|
||||||
# WebRTC
|
# WebRTC
|
||||||
from pipecat.transports.smallwebrtc.connection import SmallWebRTCConnection
|
from pipecat.transports.smallwebrtc.connection import SmallWebRTCConnection
|
||||||
from pipecat.transports.smallwebrtc.transport import SmallWebRTCTransport
|
from pipecat.transports.smallwebrtc.transport import (
|
||||||
|
SmallWebRTCOutputTransport,
|
||||||
|
SmallWebRTCTransport,
|
||||||
|
)
|
||||||
|
|
||||||
# 裸 WS 音频流
|
# 裸 WS 音频流
|
||||||
from pipecat.transports.websocket.fastapi import (
|
from pipecat.transports.websocket.fastapi import (
|
||||||
|
FastAPIWebsocketOutputTransport,
|
||||||
FastAPIWebsocketTransport,
|
FastAPIWebsocketTransport,
|
||||||
FastAPIWebsocketParams,
|
FastAPIWebsocketParams,
|
||||||
)
|
)
|
||||||
from pipecat.serializers.protobuf import ProtobufFrameSerializer
|
from pipecat.serializers.protobuf import ProtobufFrameSerializer
|
||||||
|
|
||||||
|
from services.pipecat.call_lifecycle import FixedSpeechPlaybackMarkerFrame
|
||||||
|
|
||||||
|
|
||||||
|
class _PlaybackMarkerOutputMixin:
|
||||||
|
"""Resolve fixed-speech markers after preceding audio has been sent."""
|
||||||
|
|
||||||
|
async def write_transport_frame(self, frame):
|
||||||
|
if isinstance(frame, FixedSpeechPlaybackMarkerFrame):
|
||||||
|
await frame.completion.mark_played()
|
||||||
|
return
|
||||||
|
await super().write_transport_frame(frame)
|
||||||
|
|
||||||
|
|
||||||
|
class _WebRTCOutputTransport(
|
||||||
|
_PlaybackMarkerOutputMixin,
|
||||||
|
SmallWebRTCOutputTransport,
|
||||||
|
):
|
||||||
|
pass
|
||||||
|
|
||||||
|
|
||||||
|
class _WebRTCTransport(SmallWebRTCTransport):
|
||||||
|
def output(self) -> SmallWebRTCOutputTransport:
|
||||||
|
if not self._output:
|
||||||
|
self._output = _WebRTCOutputTransport(
|
||||||
|
self._client,
|
||||||
|
self._params,
|
||||||
|
name=self._input_name,
|
||||||
|
)
|
||||||
|
return self._output
|
||||||
|
|
||||||
|
|
||||||
|
class _WebsocketOutputTransport(
|
||||||
|
_PlaybackMarkerOutputMixin,
|
||||||
|
FastAPIWebsocketOutputTransport,
|
||||||
|
):
|
||||||
|
pass
|
||||||
|
|
||||||
|
|
||||||
|
class _WebsocketTransport(FastAPIWebsocketTransport):
|
||||||
|
def __init__(self, websocket: WebSocket, params: FastAPIWebsocketParams):
|
||||||
|
super().__init__(websocket=websocket, params=params)
|
||||||
|
self._output = _WebsocketOutputTransport(
|
||||||
|
self,
|
||||||
|
self._client,
|
||||||
|
self._params,
|
||||||
|
name=self._output_name,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
def _base_params(*, video_in_enabled: bool = False) -> dict:
|
def _base_params(*, video_in_enabled: bool = False) -> dict:
|
||||||
"""两种 transport 共享的音频参数。"""
|
"""两种 transport 共享的音频参数。"""
|
||||||
@@ -41,7 +93,7 @@ def build_webrtc_transport(
|
|||||||
*,
|
*,
|
||||||
video_in_enabled: bool = False,
|
video_in_enabled: bool = False,
|
||||||
) -> SmallWebRTCTransport:
|
) -> SmallWebRTCTransport:
|
||||||
return SmallWebRTCTransport(
|
return _WebRTCTransport(
|
||||||
webrtc_connection=connection,
|
webrtc_connection=connection,
|
||||||
params=TransportParams(**_base_params(video_in_enabled=video_in_enabled)),
|
params=TransportParams(**_base_params(video_in_enabled=video_in_enabled)),
|
||||||
)
|
)
|
||||||
@@ -51,7 +103,7 @@ def build_ws_transport(websocket: WebSocket) -> FastAPIWebsocketTransport:
|
|||||||
"""裸 WS 输出。序列化用 protobuf(自定义客户端用同款解码);
|
"""裸 WS 输出。序列化用 protobuf(自定义客户端用同款解码);
|
||||||
若对接电话商,把 serializer 换成对应的 TwilioFrameSerializer 等即可。
|
若对接电话商,把 serializer 换成对应的 TwilioFrameSerializer 等即可。
|
||||||
"""
|
"""
|
||||||
return FastAPIWebsocketTransport(
|
return _WebsocketTransport(
|
||||||
websocket=websocket,
|
websocket=websocket,
|
||||||
params=FastAPIWebsocketParams(
|
params=FastAPIWebsocketParams(
|
||||||
serializer=ProtobufFrameSerializer(),
|
serializer=ProtobufFrameSerializer(),
|
||||||
|
|||||||
@@ -23,6 +23,7 @@ from services.message_stage import (
|
|||||||
MessageStageRunner,
|
MessageStageRunner,
|
||||||
MessageStageSpec,
|
MessageStageSpec,
|
||||||
)
|
)
|
||||||
|
from services.pipecat.call_lifecycle import playback_marker_for
|
||||||
from services.pipecat.realtime_tools import RealtimeTool, RealtimeToolResult
|
from services.pipecat.realtime_tools import RealtimeTool, RealtimeToolResult
|
||||||
from services.runtime_variables import DynamicVariableError, DynamicVariableStore
|
from services.runtime_variables import DynamicVariableError, DynamicVariableStore
|
||||||
from services.system_tools import state_update_properties, system_tool_kind
|
from services.system_tools import state_update_properties, system_tool_kind
|
||||||
@@ -99,6 +100,18 @@ class RealtimeWorkflowOutput(WorkflowOutput):
|
|||||||
content,
|
content,
|
||||||
suppress_transcript=True,
|
suppress_transcript=True,
|
||||||
)
|
)
|
||||||
|
playback_marker = playback_marker_for(playback_completion)
|
||||||
|
if playback_marker is not None:
|
||||||
|
async def finish_playback_tracking() -> None:
|
||||||
|
# The provider boundary is emitted after its final audio frame.
|
||||||
|
# Queueing the marker then places it behind that audio at output.
|
||||||
|
if provider_completion is not None:
|
||||||
|
await provider_completion
|
||||||
|
await self._runtime.queue_frame(playback_marker)
|
||||||
|
playback_marker.completion.mark_queued()
|
||||||
|
await playback_completion
|
||||||
|
|
||||||
|
return asyncio.create_task(finish_playback_tracking())
|
||||||
# Message playback policy and deterministic continuation must use the
|
# Message playback policy and deterministic continuation must use the
|
||||||
# transport boundary. Provider response.done only means generation
|
# transport boundary. Provider response.done only means generation
|
||||||
# has finished; audio may still be buffered at the output transport.
|
# has finished; audio may still be buffered at the output transport.
|
||||||
@@ -983,12 +996,20 @@ class WorkflowRealtimeController:
|
|||||||
return
|
return
|
||||||
self._runtime.call_end.begin("workflow_completed")
|
self._runtime.call_end.begin("workflow_completed")
|
||||||
if message:
|
if message:
|
||||||
self._runtime.call_end.arm_after_speech()
|
|
||||||
completion = await self._output.speak(
|
completion = await self._output.speak(
|
||||||
message,
|
message,
|
||||||
source="workflow-end-speech",
|
source="workflow-end-speech",
|
||||||
node_id=node_id,
|
node_id=node_id,
|
||||||
)
|
)
|
||||||
|
arm_tracked = getattr(
|
||||||
|
self._runtime.call_end,
|
||||||
|
"arm_after_tracked_speech",
|
||||||
|
None,
|
||||||
|
)
|
||||||
|
if callable(arm_tracked):
|
||||||
|
await arm_tracked()
|
||||||
|
else:
|
||||||
|
self._runtime.call_end.arm_after_speech()
|
||||||
if completion:
|
if completion:
|
||||||
await completion
|
await completion
|
||||||
else:
|
else:
|
||||||
|
|||||||
@@ -3,8 +3,15 @@ from __future__ import annotations
|
|||||||
import asyncio
|
import asyncio
|
||||||
import unittest
|
import unittest
|
||||||
|
|
||||||
from pipecat.frames.frames import BotStartedSpeakingFrame, BotStoppedSpeakingFrame
|
from pipecat.frames.frames import (
|
||||||
from services.pipecat.call_lifecycle import CallEndCoordinator
|
BotStartedSpeakingFrame,
|
||||||
|
BotStoppedSpeakingFrame,
|
||||||
|
InterruptionFrame,
|
||||||
|
)
|
||||||
|
from services.pipecat.call_lifecycle import (
|
||||||
|
CallEndCoordinator,
|
||||||
|
playback_marker_for,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
class CallEndCoordinatorTest(unittest.IsolatedAsyncioTestCase):
|
class CallEndCoordinatorTest(unittest.IsolatedAsyncioTestCase):
|
||||||
@@ -65,17 +72,68 @@ class CallEndCoordinatorTest(unittest.IsolatedAsyncioTestCase):
|
|||||||
self.coordinator.begin("workflow_completed")
|
self.coordinator.begin("workflow_completed")
|
||||||
await self.coordinator.arm_after_tracked_speech()
|
await self.coordinator.arm_after_tracked_speech()
|
||||||
|
|
||||||
await self.coordinator.observe(BotStartedSpeakingFrame())
|
first_marker = playback_marker_for(first_completion)
|
||||||
await self.coordinator.observe(BotStoppedSpeakingFrame())
|
second_marker = playback_marker_for(second_completion)
|
||||||
|
self.assertIsNotNone(first_marker)
|
||||||
|
self.assertIsNotNone(second_marker)
|
||||||
|
|
||||||
|
await first_marker.completion.mark_played()
|
||||||
self.assertTrue(first_completion.done())
|
self.assertTrue(first_completion.done())
|
||||||
self.assertFalse(second_completion.done())
|
self.assertFalse(second_completion.done())
|
||||||
self.assertEqual(self.reasons, [])
|
self.assertEqual(self.reasons, [])
|
||||||
|
|
||||||
await self.coordinator.observe(BotStartedSpeakingFrame())
|
await second_marker.completion.mark_played()
|
||||||
await self.coordinator.observe(BotStoppedSpeakingFrame())
|
|
||||||
self.assertTrue(second_completion.done())
|
self.assertTrue(second_completion.done())
|
||||||
self.assertEqual(self.reasons, ["workflow_completed"])
|
self.assertEqual(self.reasons, ["workflow_completed"])
|
||||||
|
|
||||||
|
async def test_previous_speech_stop_does_not_complete_fixed_speech(self):
|
||||||
|
await self.coordinator.observe(BotStartedSpeakingFrame())
|
||||||
|
completion = self.coordinator.track_speech()
|
||||||
|
marker = playback_marker_for(completion)
|
||||||
|
self.coordinator.begin("workflow_completed")
|
||||||
|
await self.coordinator.arm_after_tracked_speech()
|
||||||
|
|
||||||
|
await self.coordinator.observe(BotStoppedSpeakingFrame())
|
||||||
|
|
||||||
|
self.assertFalse(completion.done())
|
||||||
|
self.assertEqual(self.reasons, [])
|
||||||
|
await marker.completion.mark_played()
|
||||||
|
self.assertEqual(self.reasons, ["workflow_completed"])
|
||||||
|
|
||||||
|
async def test_delayed_previous_speech_boundary_cannot_claim_marker(self):
|
||||||
|
completion = self.coordinator.track_speech()
|
||||||
|
marker = playback_marker_for(completion)
|
||||||
|
self.coordinator.begin("workflow_completed")
|
||||||
|
await self.coordinator.arm_after_tracked_speech()
|
||||||
|
|
||||||
|
await self.coordinator.observe(BotStartedSpeakingFrame())
|
||||||
|
await self.coordinator.observe(BotStoppedSpeakingFrame())
|
||||||
|
|
||||||
|
self.assertFalse(completion.done())
|
||||||
|
self.assertEqual(self.reasons, [])
|
||||||
|
await marker.completion.mark_played()
|
||||||
|
self.assertEqual(self.reasons, ["workflow_completed"])
|
||||||
|
|
||||||
|
async def test_interruption_completes_marker_already_in_output_queue(self):
|
||||||
|
completion = self.coordinator.track_speech()
|
||||||
|
marker = playback_marker_for(completion)
|
||||||
|
marker.completion.mark_queued()
|
||||||
|
|
||||||
|
await self.coordinator.observe(InterruptionFrame())
|
||||||
|
|
||||||
|
self.assertTrue(completion.done())
|
||||||
|
|
||||||
|
async def test_interruption_does_not_complete_marker_not_yet_queued(self):
|
||||||
|
completion = self.coordinator.track_speech()
|
||||||
|
marker = playback_marker_for(completion)
|
||||||
|
|
||||||
|
await self.coordinator.observe(InterruptionFrame())
|
||||||
|
|
||||||
|
self.assertFalse(completion.done())
|
||||||
|
marker.completion.mark_queued()
|
||||||
|
await marker.completion.mark_played()
|
||||||
|
self.assertTrue(completion.done())
|
||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
unittest.main()
|
unittest.main()
|
||||||
|
|||||||
146
backend/tests/test_fixed_speech_playback.py
Normal file
146
backend/tests/test_fixed_speech_playback.py
Normal file
@@ -0,0 +1,146 @@
|
|||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import asyncio
|
||||||
|
import unittest
|
||||||
|
from types import SimpleNamespace
|
||||||
|
|
||||||
|
from models import AssistantConfig
|
||||||
|
from pipecat.frames.frames import (
|
||||||
|
BotStartedSpeakingFrame,
|
||||||
|
BotStoppedSpeakingFrame,
|
||||||
|
TTSSpeakFrame,
|
||||||
|
)
|
||||||
|
import services.brains # Initialize the brain registry before realtime imports.
|
||||||
|
from services.fixed_speech import FixedSpeechOutput
|
||||||
|
from services.pipecat.call_lifecycle import (
|
||||||
|
CallEndCoordinator,
|
||||||
|
FixedSpeechPlaybackMarkerFrame,
|
||||||
|
)
|
||||||
|
from services.pipecat.transports import build_ws_transport
|
||||||
|
from services.runtime_variables import DynamicVariableStore
|
||||||
|
from services.workflow.realtime import WorkflowRealtimeController
|
||||||
|
from services.workflow_engine import WorkflowEngine
|
||||||
|
|
||||||
|
|
||||||
|
class FixedSpeechPlaybackTest(unittest.IsolatedAsyncioTestCase):
|
||||||
|
async def test_pipeline_speech_queues_marker_immediately_after_tts(self):
|
||||||
|
queued = []
|
||||||
|
|
||||||
|
async def queue_end(_reason: str) -> None:
|
||||||
|
pass
|
||||||
|
|
||||||
|
async def queue_frame(frame) -> None:
|
||||||
|
queued.append(frame)
|
||||||
|
|
||||||
|
call_end = CallEndCoordinator(queue_end)
|
||||||
|
output = FixedSpeechOutput(
|
||||||
|
DynamicVariableStore({}),
|
||||||
|
SimpleNamespace(call_end=call_end, queue_frame=queue_frame),
|
||||||
|
)
|
||||||
|
|
||||||
|
completion = await output.speak(
|
||||||
|
"固定结束语",
|
||||||
|
source="test",
|
||||||
|
record_history=False,
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertIsInstance(queued[0], TTSSpeakFrame)
|
||||||
|
self.assertIsInstance(queued[1], FixedSpeechPlaybackMarkerFrame)
|
||||||
|
self.assertIs(queued[1].completion, completion)
|
||||||
|
|
||||||
|
async def test_websocket_output_resolves_marker(self):
|
||||||
|
async def queue_end(_reason: str) -> None:
|
||||||
|
pass
|
||||||
|
|
||||||
|
call_end = CallEndCoordinator(queue_end)
|
||||||
|
completion = call_end.track_speech()
|
||||||
|
marker = FixedSpeechPlaybackMarkerFrame(completion=completion)
|
||||||
|
websocket = SimpleNamespace(headers={})
|
||||||
|
output = build_ws_transport(websocket).output()
|
||||||
|
|
||||||
|
await output.write_transport_frame(marker)
|
||||||
|
|
||||||
|
self.assertTrue(completion.done())
|
||||||
|
|
||||||
|
async def test_realtime_end_ignores_unrelated_speech_boundaries(self):
|
||||||
|
graph = {
|
||||||
|
"specVersion": 3,
|
||||||
|
"settings": {},
|
||||||
|
"nodes": [
|
||||||
|
{"id": "start", "type": "start", "data": {}},
|
||||||
|
{
|
||||||
|
"id": "end",
|
||||||
|
"type": "end",
|
||||||
|
"data": {"message": "感谢来电,再见。", "scope": "session"},
|
||||||
|
},
|
||||||
|
],
|
||||||
|
"edges": [],
|
||||||
|
}
|
||||||
|
reasons = []
|
||||||
|
queued = []
|
||||||
|
|
||||||
|
async def queue_end(reason: str) -> None:
|
||||||
|
reasons.append(reason)
|
||||||
|
|
||||||
|
async def queue_frame(frame) -> None:
|
||||||
|
queued.append(frame)
|
||||||
|
|
||||||
|
class FakeRealtime:
|
||||||
|
def __init__(self):
|
||||||
|
self.provider_completion = None
|
||||||
|
|
||||||
|
async def update_session(self, _instructions, _tools):
|
||||||
|
pass
|
||||||
|
|
||||||
|
async def speak_fixed(self, _text, *, suppress_transcript=True):
|
||||||
|
self.provider_completion = (
|
||||||
|
asyncio.get_running_loop().create_future()
|
||||||
|
)
|
||||||
|
return self.provider_completion
|
||||||
|
|
||||||
|
call_end = CallEndCoordinator(queue_end)
|
||||||
|
realtime = FakeRealtime()
|
||||||
|
controller = WorkflowRealtimeController(
|
||||||
|
cfg=AssistantConfig(type="workflow", graph=graph),
|
||||||
|
engine=WorkflowEngine(graph),
|
||||||
|
store=DynamicVariableStore({}),
|
||||||
|
runtime=SimpleNamespace(
|
||||||
|
realtime=realtime,
|
||||||
|
queue_frame=queue_frame,
|
||||||
|
call_end=call_end,
|
||||||
|
session_id="test-session",
|
||||||
|
client_tools=None,
|
||||||
|
set_input_enabled=lambda _enabled: None,
|
||||||
|
capture_image=None,
|
||||||
|
),
|
||||||
|
)
|
||||||
|
|
||||||
|
end_task = asyncio.create_task(controller._enter_end("end"))
|
||||||
|
while realtime.provider_completion is None:
|
||||||
|
await asyncio.sleep(0)
|
||||||
|
|
||||||
|
await call_end.observe(BotStartedSpeakingFrame())
|
||||||
|
await call_end.observe(BotStoppedSpeakingFrame())
|
||||||
|
self.assertEqual(reasons, [])
|
||||||
|
self.assertFalse(end_task.done())
|
||||||
|
|
||||||
|
realtime.provider_completion.set_result(None)
|
||||||
|
marker = None
|
||||||
|
while marker is None:
|
||||||
|
await asyncio.sleep(0)
|
||||||
|
marker = next(
|
||||||
|
(
|
||||||
|
frame
|
||||||
|
for frame in queued
|
||||||
|
if isinstance(frame, FixedSpeechPlaybackMarkerFrame)
|
||||||
|
),
|
||||||
|
None,
|
||||||
|
)
|
||||||
|
await marker.completion.mark_played()
|
||||||
|
await end_task
|
||||||
|
|
||||||
|
self.assertEqual(reasons, ["workflow_completed"])
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
unittest.main()
|
||||||
Reference in New Issue
Block a user