Files
ai-video-fullstack/backend/tests/test_fixed_speech_playback.py

147 lines
4.6 KiB
Python

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()