777 lines
32 KiB
Python
777 lines
32 KiB
Python
from __future__ import annotations
|
|
|
|
import time
|
|
import unittest
|
|
|
|
from owner_voice_pet.config import AppConfig
|
|
from owner_voice_pet.conversation import ConversationContext
|
|
from owner_voice_pet.events import (
|
|
ACK_STARTED,
|
|
BARGE_IN_DETECTED,
|
|
CAPTURE_STARTED,
|
|
CONTINUATION_DECISION_MADE,
|
|
FOLLOWUP_LISTENING,
|
|
FOLLOWUP_TIMEOUT,
|
|
LLM_STARTED,
|
|
PLAYBACK_FINISHED,
|
|
PLAYBACK_INTERRUPTED,
|
|
SPEECH_ENDED,
|
|
SPEECH_STARTED,
|
|
STANDBY_RESUMED,
|
|
STT_STARTED,
|
|
TRANSCRIPT_FINAL,
|
|
TRANSCRIPT_PARTIAL,
|
|
TTS_STARTED,
|
|
WAKE_DETECTED,
|
|
WAKE_LISTENING,
|
|
PipelineEventBus,
|
|
)
|
|
from owner_voice_pet.llm import MockLlmProvider
|
|
from owner_voice_pet.models import AudioFrame, AudioSegment, ErrorCode, PlaybackResult, ReplyDelta, Transcript, TransportHealth
|
|
from owner_voice_pet.assistant_pipeline import VoiceAssistantPipeline
|
|
from owner_voice_pet.runtime import build_live_runtime
|
|
from owner_voice_pet.stt import MetadataSttProvider, SherpaOnnxSttProvider
|
|
from owner_voice_pet.transport import MemoryAudioTransport
|
|
from owner_voice_pet.tts import SineTtsProvider
|
|
from owner_voice_pet.vad import EnergyVadProvider, PrimarySpeakerVadRecorder, VadRecorder
|
|
from owner_voice_pet.wakeword import KeywordWakeWordProvider
|
|
|
|
|
|
def segment_frames(start_id: int, start_ms: int, partials: list[str] | None = None) -> list[AudioFrame]:
|
|
partials = partials or []
|
|
first_metadata: dict[str, object] = {"duration_ms": 20, "speech": True}
|
|
second_metadata: dict[str, object] = {"duration_ms": 20, "speech": True}
|
|
if len(partials) >= 1:
|
|
first_metadata["partial_transcript"] = partials[0]
|
|
if len(partials) >= 2:
|
|
second_metadata["partial_transcript"] = partials[1]
|
|
return [
|
|
AudioFrame(b"\xff\x7f", 16000, 1, start_ms, start_id, first_metadata),
|
|
AudioFrame(b"\xff\x7f", 16000, 1, start_ms + 20, start_id + 1, second_metadata),
|
|
AudioFrame(b"\x00\x00", 16000, 1, start_ms + 40, start_id + 2, {"duration_ms": 20, "speech": False}),
|
|
AudioFrame(b"\x00\x00", 16000, 1, start_ms + 60, start_id + 3, {"duration_ms": 20, "speech": False}),
|
|
]
|
|
|
|
|
|
def long_speech_frames_with_stale_partial(start_id: int, start_ms: int, duration_ms: int) -> list[AudioFrame]:
|
|
frames: list[AudioFrame] = []
|
|
for index in range(duration_ms // 20):
|
|
metadata: dict[str, object] = {
|
|
"duration_ms": 20,
|
|
"speech": True,
|
|
"partial_transcript": "你知道",
|
|
"transcript": "你知道我在说什么吗",
|
|
}
|
|
frames.append(
|
|
AudioFrame(
|
|
b"\xff\x7f",
|
|
16000,
|
|
1,
|
|
start_ms + index * 20,
|
|
start_id + index,
|
|
metadata,
|
|
)
|
|
)
|
|
return frames
|
|
|
|
|
|
def wake_frame(frame_id: int, timestamp_ms: int) -> AudioFrame:
|
|
return AudioFrame(
|
|
b"\xff\x7f",
|
|
16000,
|
|
1,
|
|
timestamp_ms,
|
|
frame_id,
|
|
{"duration_ms": 20, "wake_word": "小杰小杰", "wake_confidence": 0.95},
|
|
)
|
|
|
|
|
|
class QueueSttProvider:
|
|
def __init__(self, texts: list[str]) -> None:
|
|
self.texts = list(texts)
|
|
self.calls: list[AudioSegment] = []
|
|
self.loaded = False
|
|
|
|
def load(self) -> None:
|
|
self.loaded = True
|
|
|
|
def transcribe(self, segment: AudioSegment) -> Transcript:
|
|
self.calls.append(segment)
|
|
text = self.texts.pop(0)
|
|
return Transcript(text, "zh", 1.0, segment.duration_ms, "queue-stt")
|
|
|
|
|
|
class QueueLlmProvider:
|
|
def __init__(self, replies: list[list[str]]) -> None:
|
|
self.replies = [list(item) for item in replies]
|
|
self.calls = []
|
|
|
|
def stream_reply(self, messages):
|
|
self.calls.append(list(messages))
|
|
chunks = self.replies.pop(0)
|
|
for chunk in chunks:
|
|
yield ReplyDelta(chunk)
|
|
yield ReplyDelta("", finish_reason="stop")
|
|
|
|
|
|
class CountingTtsProvider:
|
|
def __init__(self) -> None:
|
|
self.delegate = SineTtsProvider()
|
|
self.synthesized_texts: list[str] = []
|
|
|
|
def load(self) -> None:
|
|
self.delegate.load()
|
|
|
|
def synthesize(self, text: str) -> AudioSegment:
|
|
self.synthesized_texts.append(text)
|
|
return self.delegate.synthesize(text)
|
|
|
|
|
|
class RecordingReporter:
|
|
def __init__(self) -> None:
|
|
self.statuses: list[str] = []
|
|
self.transcripts: list[str] = []
|
|
self.partials: list[str] = []
|
|
self.errors: list[str] = []
|
|
self.events: list[str] = []
|
|
|
|
def status(self, state: str, message: str, *, turn_id: int | None = None) -> None:
|
|
self.statuses.append(message)
|
|
self.events.append(f"status:{message}")
|
|
|
|
def transcript(self, text: str, *, final: bool, turn_id: int | None = None) -> None:
|
|
if final:
|
|
self.transcripts.append(text)
|
|
self.events.append(f"transcript:final:{text}")
|
|
else:
|
|
self.partials.append(text)
|
|
self.events.append(f"transcript:partial:{text}")
|
|
|
|
def error(self, stage: str, code: str, message: str, *, turn_id: int | None = None) -> None:
|
|
self.errors.append(f"{stage}:{code}:{message}")
|
|
|
|
|
|
class MarkerAudioPreprocessor:
|
|
def __init__(self, partial_text: str = "降噪后问题") -> None:
|
|
self.partial_text = partial_text
|
|
self.loaded = False
|
|
self.reset_calls = 0
|
|
self.frames: list[AudioFrame] = []
|
|
|
|
def load(self) -> None:
|
|
self.loaded = True
|
|
|
|
def reset(self) -> None:
|
|
self.reset_calls += 1
|
|
|
|
def process_frame(self, frame: AudioFrame) -> AudioFrame:
|
|
metadata = dict(frame.metadata)
|
|
metadata["denoised"] = True
|
|
metadata["partial_transcript"] = self.partial_text
|
|
processed = AudioFrame(
|
|
b"\x01\x00",
|
|
frame.sample_rate,
|
|
frame.channels,
|
|
frame.timestamp_ms,
|
|
frame.frame_id,
|
|
metadata,
|
|
)
|
|
self.frames.append(processed)
|
|
return processed
|
|
|
|
def flush(self) -> list[AudioFrame]:
|
|
return []
|
|
|
|
|
|
def make_runtime(
|
|
texts: list[str],
|
|
context: ConversationContext | None = None,
|
|
partial_texts: list[list[str]] | None = None,
|
|
audio_preprocessor: MarkerAudioPreprocessor | None = None,
|
|
wake_ack_text: str = "我在",
|
|
) -> tuple[VoiceAssistantPipeline, QueueSttProvider, MockLlmProvider, MemoryAudioTransport, RecordingReporter]:
|
|
frames = []
|
|
for idx, _text in enumerate(texts):
|
|
base_id = idx * 5
|
|
base_ms = idx * 120
|
|
frames.append(wake_frame(base_id, base_ms))
|
|
partials = partial_texts[idx] if partial_texts and idx < len(partial_texts) else None
|
|
frames.extend(segment_frames(base_id + 1, base_ms + 20, partials=partials))
|
|
transport = MemoryAudioTransport(frames, flush_clears_input=False)
|
|
stt = QueueSttProvider(texts)
|
|
llm = MockLlmProvider(["这是答复。"])
|
|
tts = SineTtsProvider()
|
|
reporter = RecordingReporter()
|
|
event_bus = PipelineEventBus()
|
|
runtime = VoiceAssistantPipeline(
|
|
config=AppConfig(llm_api_key="secret", speech_provider="cloud", wake_ack_text=wake_ack_text),
|
|
transport=transport,
|
|
wakeword=KeywordWakeWordProvider(),
|
|
vad_recorder=VadRecorder(EnergyVadProvider(), min_duration_ms=40, end_silence_ms=40),
|
|
stt=stt,
|
|
audio_preprocessor=audio_preprocessor,
|
|
realtime_stt=MetadataSttProvider() if partial_texts is not None else None,
|
|
llm=llm,
|
|
tts=tts,
|
|
context=context or ConversationContext(),
|
|
reporter=reporter,
|
|
event_bus=event_bus,
|
|
)
|
|
return runtime, stt, llm, transport, reporter
|
|
|
|
|
|
class PlaybackInjectedTransport(MemoryAudioTransport):
|
|
def __init__(
|
|
self,
|
|
frames: list[AudioFrame],
|
|
*,
|
|
inject_after_play_count: int,
|
|
injected_frames: list[AudioFrame],
|
|
) -> None:
|
|
super().__init__(frames, flush_clears_input=False)
|
|
self.inject_after_play_count = inject_after_play_count
|
|
self.injected_frames = list(injected_frames)
|
|
self.play_count = 0
|
|
|
|
def play_pcm(self, segment: AudioSegment, interrupt: bool = False) -> PlaybackResult:
|
|
result = super().play_pcm(segment, interrupt=interrupt)
|
|
self.play_count += 1
|
|
if self.play_count == self.inject_after_play_count:
|
|
for frame in self.injected_frames:
|
|
self.inject(frame)
|
|
time.sleep(0.002)
|
|
return result
|
|
|
|
def health(self) -> TransportHealth:
|
|
return TransportHealth(True, True, "playback injected")
|
|
|
|
|
|
def silence_frames(start_id: int, start_ms: int, count: int) -> list[AudioFrame]:
|
|
return [
|
|
AudioFrame(
|
|
b"\x00\x00",
|
|
16000,
|
|
1,
|
|
start_ms + index * 20,
|
|
start_id + index,
|
|
{"duration_ms": 20, "speech": False},
|
|
)
|
|
for index in range(count)
|
|
]
|
|
|
|
|
|
class LiveRuntimeTests(unittest.TestCase):
|
|
def test_agent_runtime_listens_without_wake_word(self) -> None:
|
|
frames = segment_frames(1, 20, partials=["直接提问", "直接提问"])
|
|
transport = MemoryAudioTransport(frames, flush_clears_input=False)
|
|
stt = QueueSttProvider(["直接提问"])
|
|
llm = QueueLlmProvider([["这是全双工回答。"]])
|
|
reporter = RecordingReporter()
|
|
runtime = VoiceAssistantPipeline(
|
|
config=AppConfig(
|
|
assistant_mode="full_duplex_agent",
|
|
llm_api_key="secret",
|
|
speech_provider="cloud",
|
|
wake_ack_text="",
|
|
),
|
|
transport=transport,
|
|
wakeword=KeywordWakeWordProvider(),
|
|
vad_recorder=VadRecorder(EnergyVadProvider(), min_duration_ms=40, end_silence_ms=40),
|
|
stt=stt,
|
|
realtime_stt=MetadataSttProvider(),
|
|
llm=llm,
|
|
tts=SineTtsProvider(),
|
|
context=ConversationContext(),
|
|
reporter=reporter,
|
|
event_bus=PipelineEventBus(),
|
|
)
|
|
|
|
summary = runtime.run_agent(once=True)
|
|
|
|
self.assertEqual(summary.completed_turns, 1)
|
|
self.assertEqual(stt.calls[0].metadata["end_reason"], "silence")
|
|
self.assertEqual(reporter.transcripts, ["直接提问"])
|
|
self.assertIn("监听中:请直接说话", reporter.statuses)
|
|
self.assertNotIn("唤醒命中", reporter.statuses)
|
|
self.assertIn("恢复监听:可直接说话", reporter.statuses)
|
|
|
|
def test_repeated_runtime_runs_two_turns_and_returns_to_standby(self) -> None:
|
|
runtime, stt, llm, transport, reporter = make_runtime(["第一问", "第二问"])
|
|
self.assertIsInstance(runtime, VoiceAssistantPipeline)
|
|
self.assertIsNotNone(runtime.controller)
|
|
summary = runtime.run(max_turns=2)
|
|
|
|
self.assertEqual(summary.completed_turns, 2)
|
|
self.assertEqual(runtime.config.post_playback_drain_ms, 0)
|
|
self.assertEqual(len(stt.calls), 2)
|
|
self.assertEqual(len(llm.calls), 2)
|
|
self.assertEqual(len(transport.played_segments), 6)
|
|
self.assertEqual(transport.flush_count, 6)
|
|
self.assertEqual(transport.played_segments[2].metadata["chime"], "end")
|
|
self.assertEqual(transport.played_segments[2].metadata["source"], "file")
|
|
self.assertEqual(transport.played_segments[5].metadata["chime"], "end")
|
|
self.assertEqual(transport.played_segments[5].metadata["source"], "file")
|
|
self.assertEqual(reporter.transcripts, ["第一问", "第二问"])
|
|
self.assertIn("应答中:我在", reporter.statuses)
|
|
self.assertLess(reporter.statuses.index("唤醒命中"), reporter.statuses.index("应答中:我在"))
|
|
self.assertLess(reporter.statuses.index("应答中:我在"), reporter.statuses.index("请说出问题"))
|
|
self.assertLess(reporter.statuses.index("请说出问题"), reporter.statuses.index("录音中:正在听取问题"))
|
|
self.assertIn("恢复待机:可继续唤醒", reporter.statuses[-1])
|
|
event_types = [event.type for event in runtime.event_bus.events]
|
|
expected_order = [
|
|
WAKE_LISTENING,
|
|
WAKE_DETECTED,
|
|
ACK_STARTED,
|
|
CAPTURE_STARTED,
|
|
SPEECH_STARTED,
|
|
SPEECH_ENDED,
|
|
STT_STARTED,
|
|
TRANSCRIPT_FINAL,
|
|
LLM_STARTED,
|
|
TTS_STARTED,
|
|
PLAYBACK_FINISHED,
|
|
STANDBY_RESUMED,
|
|
]
|
|
positions = [event_types.index(item) for item in expected_order]
|
|
self.assertEqual(positions, sorted(positions))
|
|
|
|
def test_wake_ack_audio_is_prepared_once_and_reused(self) -> None:
|
|
frames = []
|
|
for idx, _text in enumerate(["第一问", "第二问"]):
|
|
base_id = idx * 5
|
|
base_ms = idx * 120
|
|
frames.append(wake_frame(base_id, base_ms))
|
|
frames.extend(segment_frames(base_id + 1, base_ms + 20))
|
|
transport = MemoryAudioTransport(frames, flush_clears_input=False)
|
|
ack_tts = CountingTtsProvider()
|
|
runtime = VoiceAssistantPipeline(
|
|
config=AppConfig(llm_api_key="secret", speech_provider="cloud", wake_ack_text="我在"),
|
|
transport=transport,
|
|
wakeword=KeywordWakeWordProvider(),
|
|
vad_recorder=VadRecorder(EnergyVadProvider(), min_duration_ms=40, end_silence_ms=40),
|
|
stt=QueueSttProvider(["第一问", "第二问"]),
|
|
realtime_stt=None,
|
|
llm=MockLlmProvider(["这是答复。"]),
|
|
tts=SineTtsProvider(),
|
|
ack_tts=ack_tts,
|
|
context=ConversationContext(),
|
|
reporter=RecordingReporter(),
|
|
event_bus=PipelineEventBus(),
|
|
)
|
|
|
|
summary = runtime.run(max_turns=2)
|
|
|
|
self.assertEqual(summary.completed_turns, 2)
|
|
self.assertEqual(ack_tts.synthesized_texts, ["我在"])
|
|
self.assertEqual(transport.played_segments[0].metadata["text"], "我在")
|
|
self.assertEqual(transport.played_segments[3].metadata["text"], "我在")
|
|
|
|
def test_zero_post_playback_drain_flushes_without_dropping_prefilled_question(self) -> None:
|
|
runtime, stt, _, transport, reporter = make_runtime(["第一问"])
|
|
self.assertEqual(runtime.config.post_playback_drain_ms, 0)
|
|
|
|
summary = runtime.run(max_turns=1)
|
|
|
|
self.assertEqual(summary.completed_turns, 1)
|
|
self.assertEqual(reporter.transcripts, ["第一问"])
|
|
self.assertEqual(len(stt.calls), 1)
|
|
self.assertEqual(transport.flush_count, 3)
|
|
self.assertEqual(transport.played_segments[-1].metadata["chime"], "end")
|
|
self.assertEqual(transport.played_segments[-1].metadata["source"], "file")
|
|
|
|
def test_no_ack_text_does_not_drain_before_capture(self) -> None:
|
|
runtime, stt, _, transport, reporter = make_runtime(["第一问"], wake_ack_text="")
|
|
|
|
summary = runtime.run(max_turns=1)
|
|
|
|
self.assertEqual(summary.completed_turns, 1)
|
|
self.assertEqual(reporter.transcripts, ["第一问"])
|
|
self.assertEqual(len(stt.calls), 1)
|
|
self.assertEqual(len(transport.played_segments), 2)
|
|
self.assertEqual(transport.flush_count, 2)
|
|
self.assertEqual(transport.played_segments[-1].metadata["chime"], "end")
|
|
self.assertEqual(transport.played_segments[-1].metadata["source"], "file")
|
|
|
|
def test_end_chime_can_be_disabled(self) -> None:
|
|
runtime, _, _, transport, _ = make_runtime(["第一问"], wake_ack_text="")
|
|
runtime.config = AppConfig(llm_api_key="secret", speech_provider="cloud", wake_ack_text="", end_chime_enabled=False)
|
|
runtime.controller.config = runtime.config
|
|
|
|
summary = runtime.run(max_turns=1)
|
|
|
|
self.assertEqual(summary.completed_turns, 1)
|
|
self.assertEqual(len(transport.played_segments), 1)
|
|
self.assertNotIn("chime", transport.played_segments[-1].metadata)
|
|
|
|
def test_temporary_context_is_sent_to_second_llm_call(self) -> None:
|
|
runtime, _, llm, _, _ = make_runtime(["第一问", "第二问"])
|
|
runtime.run(max_turns=2)
|
|
|
|
second_call_text = [message.content for message in llm.calls[1]]
|
|
self.assertIn("第一问", second_call_text)
|
|
self.assertIn("这是答复。", second_call_text)
|
|
self.assertEqual(second_call_text[-1], "第二问")
|
|
|
|
def test_new_runtime_context_starts_empty(self) -> None:
|
|
first_context = ConversationContext()
|
|
first_runtime, _, _, _, _ = make_runtime(["第一问"], context=first_context)
|
|
first_runtime.run(max_turns=1)
|
|
self.assertGreater(len(first_context.messages()), 0)
|
|
|
|
second_context = ConversationContext()
|
|
make_runtime(["第二问"], context=second_context)
|
|
self.assertEqual(second_context.messages(), ())
|
|
|
|
def test_transcript_is_reported_before_llm_thinking(self) -> None:
|
|
runtime, _, _, _, reporter = make_runtime(["第一问"])
|
|
runtime.run(max_turns=1)
|
|
|
|
transcript_index = reporter.events.index("transcript:final:第一问")
|
|
thinking_index = next(
|
|
index for index, event in enumerate(reporter.events) if event == "status:思考中:正在生成回复"
|
|
)
|
|
self.assertLess(transcript_index, thinking_index)
|
|
|
|
def test_realtime_transcript_is_reported_while_capturing(self) -> None:
|
|
runtime, _, llm, _, reporter = make_runtime(["第一问"], partial_texts=[["第一", "第一问"]])
|
|
runtime.run(max_turns=1)
|
|
|
|
self.assertEqual(reporter.partials, ["第一问"])
|
|
self.assertEqual(reporter.transcripts, ["第一问"])
|
|
self.assertEqual(llm.calls[0][-1].content, "第一问")
|
|
event_types = [event.type for event in runtime.event_bus.events]
|
|
self.assertLess(event_types.index(SPEECH_STARTED), event_types.index(TRANSCRIPT_PARTIAL))
|
|
self.assertLess(event_types.index(TRANSCRIPT_PARTIAL), event_types.index(SPEECH_ENDED))
|
|
self.assertLess(event_types.index(TRANSCRIPT_PARTIAL), event_types.index(TRANSCRIPT_FINAL))
|
|
|
|
def test_realtime_transcript_idle_ends_current_utterance(self) -> None:
|
|
frames = [wake_frame(0, 0)]
|
|
frames.extend(long_speech_frames_with_stale_partial(1, 20, 2200))
|
|
transport = MemoryAudioTransport(frames, flush_clears_input=False)
|
|
stt = QueueSttProvider(["你知道我在说什么吗"])
|
|
llm = MockLlmProvider(["这是答复。"])
|
|
reporter = RecordingReporter()
|
|
runtime = VoiceAssistantPipeline(
|
|
config=AppConfig(
|
|
llm_api_key="secret",
|
|
speech_provider="cloud",
|
|
wake_ack_text="",
|
|
realtime_transcript_idle_timeout_ms=1500,
|
|
),
|
|
transport=transport,
|
|
wakeword=KeywordWakeWordProvider(),
|
|
vad_recorder=VadRecorder(
|
|
EnergyVadProvider(),
|
|
min_duration_ms=40,
|
|
end_silence_ms=5000,
|
|
max_recording_ms=10000,
|
|
),
|
|
stt=stt,
|
|
audio_preprocessor=MarkerAudioPreprocessor(partial_text="你知道"),
|
|
realtime_stt=MetadataSttProvider(),
|
|
llm=llm,
|
|
tts=SineTtsProvider(),
|
|
context=ConversationContext(),
|
|
reporter=reporter,
|
|
event_bus=PipelineEventBus(),
|
|
)
|
|
|
|
summary = runtime.run(max_turns=1)
|
|
|
|
self.assertEqual(summary.completed_turns, 1)
|
|
self.assertEqual(reporter.partials, ["你知道"])
|
|
self.assertEqual(reporter.transcripts, ["你知道我在说什么吗"])
|
|
self.assertEqual(len(stt.calls), 1)
|
|
self.assertEqual(stt.calls[0].metadata["end_reason"], "partial_transcript_idle")
|
|
self.assertLess(stt.calls[0].duration_ms, 1700)
|
|
|
|
def test_capture_uses_denoised_frames_for_partial_and_final_stt(self) -> None:
|
|
preprocessor = MarkerAudioPreprocessor(partial_text="降噪后问题")
|
|
runtime, stt, _, _, reporter = make_runtime(
|
|
["第一问"],
|
|
partial_texts=[["原始噪声", "原始噪声"]],
|
|
audio_preprocessor=preprocessor,
|
|
)
|
|
runtime.run(max_turns=1)
|
|
|
|
self.assertTrue(preprocessor.loaded)
|
|
self.assertGreaterEqual(preprocessor.reset_calls, 1)
|
|
self.assertEqual(reporter.partials, ["降噪后问题"])
|
|
self.assertEqual(len(stt.calls), 1)
|
|
self.assertTrue(stt.calls[0].metadata["denoised"])
|
|
self.assertIn(b"\x01\x00", stt.calls[0].pcm)
|
|
|
|
def test_wake_keyword_does_not_pollute_llm_user_message(self) -> None:
|
|
runtime, _, llm, _, _ = make_runtime(["第一问"])
|
|
runtime.run(max_turns=1)
|
|
|
|
self.assertEqual(llm.calls[0][-1].content, "第一问")
|
|
self.assertNotIn("小杰小杰", llm.calls[0][-1].content)
|
|
|
|
def test_assistant_reply_sanitizes_tts_text_and_context(self) -> None:
|
|
frames = [wake_frame(0, 0)]
|
|
frames.extend(segment_frames(1, 20, partials=["第一问", "第一问"]))
|
|
transport = MemoryAudioTransport(frames, flush_clears_input=False)
|
|
stt = QueueSttProvider(["第一问"])
|
|
llm = QueueLlmProvider([["你好 😊。没问题[捂脸],我来帮你。"]])
|
|
context = ConversationContext()
|
|
runtime = VoiceAssistantPipeline(
|
|
config=AppConfig(llm_api_key="secret", speech_provider="cloud", wake_ack_text=""),
|
|
transport=transport,
|
|
wakeword=KeywordWakeWordProvider(),
|
|
vad_recorder=VadRecorder(EnergyVadProvider(), min_duration_ms=40, end_silence_ms=40),
|
|
stt=stt,
|
|
realtime_stt=MetadataSttProvider(),
|
|
llm=llm,
|
|
tts=SineTtsProvider(),
|
|
context=context,
|
|
reporter=RecordingReporter(),
|
|
event_bus=PipelineEventBus(),
|
|
)
|
|
|
|
summary = runtime.run(max_turns=1)
|
|
context_texts = [message.content for message in context.messages()]
|
|
|
|
self.assertEqual(summary.completed_turns, 1)
|
|
self.assertEqual(transport.played_segments[0].metadata["text"], "你好。没问题,我来帮你。")
|
|
self.assertIn("你好。没问题,我来帮你。", context_texts)
|
|
self.assertFalse(any("😊" in text or "[捂脸]" in text for text in context_texts))
|
|
|
|
def test_emoji_only_reply_recovers_without_tts_playback(self) -> None:
|
|
frames = [wake_frame(0, 0)]
|
|
frames.extend(segment_frames(1, 20, partials=["第一问", "第一问"]))
|
|
transport = MemoryAudioTransport(frames, flush_clears_input=False)
|
|
runtime = VoiceAssistantPipeline(
|
|
config=AppConfig(llm_api_key="secret", speech_provider="cloud", wake_ack_text=""),
|
|
transport=transport,
|
|
wakeword=KeywordWakeWordProvider(),
|
|
vad_recorder=VadRecorder(EnergyVadProvider(), min_duration_ms=40, end_silence_ms=40),
|
|
stt=QueueSttProvider(["第一问"]),
|
|
realtime_stt=MetadataSttProvider(),
|
|
llm=QueueLlmProvider([["😂😂"]]),
|
|
tts=SineTtsProvider(),
|
|
context=ConversationContext(),
|
|
reporter=RecordingReporter(),
|
|
event_bus=PipelineEventBus(),
|
|
)
|
|
|
|
summary = runtime.run(once=True)
|
|
event_types = [event.type for event in runtime.event_bus.events]
|
|
|
|
self.assertEqual(summary.completed_turns, 0)
|
|
self.assertEqual(summary.failed_turns, 1)
|
|
self.assertIsNotNone(summary.last_error)
|
|
self.assertEqual(summary.last_error.code, ErrorCode.TTS_EMPTY_AUDIO)
|
|
self.assertEqual(transport.played_segments, [])
|
|
self.assertNotIn(TTS_STARTED, event_types)
|
|
self.assertEqual(event_types[-1], STANDBY_RESUMED)
|
|
|
|
def test_live_runtime_uses_primary_speaker_endpoint_by_default(self) -> None:
|
|
runtime = build_live_runtime(AppConfig(llm_api_key="secret"))
|
|
self.assertIsInstance(runtime.vad_recorder, PrimarySpeakerVadRecorder)
|
|
self.assertEqual(runtime.config.speech_provider, "local")
|
|
self.assertIsInstance(runtime.stt, SherpaOnnxSttProvider)
|
|
self.assertIsNotNone(runtime.realtime_stt)
|
|
|
|
def test_assistant_followup_question_enters_listening_without_second_wake(self) -> None:
|
|
frames = [wake_frame(0, 0)]
|
|
frames.extend(segment_frames(1, 20, partials=["第一问", "第一问"]))
|
|
frames.extend(segment_frames(10, 400, partials=["继续内容", "继续内容"]))
|
|
transport = MemoryAudioTransport(frames, flush_clears_input=False)
|
|
stt = QueueSttProvider(["第一问", "继续内容"])
|
|
llm = QueueLlmProvider([["你想继续听哪一部分?"], ["这是补充回答。"]])
|
|
reporter = RecordingReporter()
|
|
runtime = VoiceAssistantPipeline(
|
|
config=AppConfig(
|
|
llm_api_key="secret",
|
|
speech_provider="cloud",
|
|
wake_ack_text="",
|
|
followup_listen_timeout_ms=3000,
|
|
),
|
|
transport=transport,
|
|
wakeword=KeywordWakeWordProvider(),
|
|
vad_recorder=VadRecorder(EnergyVadProvider(), min_duration_ms=40, end_silence_ms=40),
|
|
stt=stt,
|
|
realtime_stt=MetadataSttProvider(),
|
|
llm=llm,
|
|
tts=SineTtsProvider(),
|
|
context=ConversationContext(),
|
|
reporter=reporter,
|
|
event_bus=PipelineEventBus(),
|
|
)
|
|
|
|
summary = runtime.run(max_turns=1)
|
|
event_types = [event.type for event in runtime.event_bus.events]
|
|
|
|
self.assertEqual(summary.completed_turns, 2)
|
|
self.assertEqual(event_types.count(WAKE_DETECTED), 1)
|
|
self.assertEqual(event_types.count(FOLLOWUP_LISTENING), 1)
|
|
self.assertEqual(event_types.count(STT_STARTED), 2)
|
|
self.assertEqual(len(llm.calls), 2)
|
|
self.assertEqual(reporter.transcripts, ["第一问", "继续内容"])
|
|
second_call_text = [message.content for message in llm.calls[1]]
|
|
self.assertIn("第一问", second_call_text)
|
|
self.assertIn("你想继续听哪一部分?", second_call_text)
|
|
|
|
def test_followup_timeout_returns_to_standby(self) -> None:
|
|
frames = [wake_frame(0, 0)]
|
|
frames.extend(segment_frames(1, 20, partials=["第一问", "第一问"]))
|
|
frames.extend(silence_frames(20, 400, 170))
|
|
transport = MemoryAudioTransport(frames, flush_clears_input=False)
|
|
stt = QueueSttProvider(["第一问"])
|
|
llm = QueueLlmProvider([["你想继续听哪一部分?"]])
|
|
reporter = RecordingReporter()
|
|
runtime = VoiceAssistantPipeline(
|
|
config=AppConfig(
|
|
llm_api_key="secret",
|
|
speech_provider="cloud",
|
|
wake_ack_text="",
|
|
followup_listen_timeout_ms=3000,
|
|
),
|
|
transport=transport,
|
|
wakeword=KeywordWakeWordProvider(),
|
|
vad_recorder=VadRecorder(EnergyVadProvider(), min_duration_ms=40, end_silence_ms=40),
|
|
stt=stt,
|
|
realtime_stt=MetadataSttProvider(),
|
|
llm=llm,
|
|
tts=SineTtsProvider(),
|
|
context=ConversationContext(),
|
|
reporter=reporter,
|
|
event_bus=PipelineEventBus(),
|
|
)
|
|
|
|
summary = runtime.run(max_turns=1)
|
|
event_types = [event.type for event in runtime.event_bus.events]
|
|
|
|
self.assertEqual(summary.completed_turns, 1)
|
|
self.assertEqual(summary.failed_turns, 0)
|
|
self.assertIn(FOLLOWUP_TIMEOUT, event_types)
|
|
self.assertEqual(event_types[-1], STANDBY_RESUMED)
|
|
self.assertEqual(len(llm.calls), 1)
|
|
self.assertEqual(transport.played_segments[-1].metadata["chime"], "end")
|
|
self.assertEqual(transport.played_segments[-1].metadata["source"], "file")
|
|
|
|
def test_completed_reply_returns_to_standby_without_cloud_classifier_delay(self) -> None:
|
|
frames = [wake_frame(0, 0)]
|
|
frames.extend(segment_frames(1, 20, partials=["没有呢", "没有呢"]))
|
|
transport = MemoryAudioTransport(frames, flush_clears_input=False)
|
|
llm = QueueLlmProvider([["明白了,我先保持待机。有需要再叫我就行。"]])
|
|
runtime = VoiceAssistantPipeline(
|
|
config=AppConfig(
|
|
llm_api_key="secret",
|
|
speech_provider="cloud",
|
|
wake_ack_text="",
|
|
followup_listen_timeout_ms=3000,
|
|
),
|
|
transport=transport,
|
|
wakeword=KeywordWakeWordProvider(),
|
|
vad_recorder=VadRecorder(EnergyVadProvider(), min_duration_ms=40, end_silence_ms=40),
|
|
stt=QueueSttProvider(["没有呢"]),
|
|
realtime_stt=MetadataSttProvider(),
|
|
llm=llm,
|
|
tts=SineTtsProvider(),
|
|
context=ConversationContext(),
|
|
reporter=RecordingReporter(),
|
|
event_bus=PipelineEventBus(),
|
|
)
|
|
|
|
summary = runtime.run(max_turns=1)
|
|
event_types = [event.type for event in runtime.event_bus.events]
|
|
|
|
self.assertEqual(summary.completed_turns, 1)
|
|
self.assertNotIn(FOLLOWUP_LISTENING, event_types)
|
|
self.assertEqual(event_types[-1], STANDBY_RESUMED)
|
|
self.assertEqual(len(llm.calls), 1)
|
|
|
|
def test_barge_in_interrupts_playback_and_keeps_only_spoken_context(self) -> None:
|
|
first_question_frames = [wake_frame(0, 0)]
|
|
first_question_frames.extend(
|
|
[
|
|
AudioFrame(
|
|
frame.pcm,
|
|
frame.sample_rate,
|
|
frame.channels,
|
|
frame.timestamp_ms,
|
|
frame.frame_id,
|
|
{**dict(frame.metadata), "speaker_id": "owner"},
|
|
)
|
|
for frame in segment_frames(1, 20, partials=["第一问", "第一问"])
|
|
]
|
|
)
|
|
barge_frames = [
|
|
AudioFrame(
|
|
b"\xff\x7f",
|
|
16000,
|
|
1,
|
|
600,
|
|
20,
|
|
{"duration_ms": 20, "speech": True, "partial_transcript": "等一下", "speaker_id": "owner"},
|
|
),
|
|
AudioFrame(
|
|
b"\xff\x7f",
|
|
16000,
|
|
1,
|
|
620,
|
|
21,
|
|
{"duration_ms": 20, "speech": True, "partial_transcript": "等一下", "speaker_id": "owner"},
|
|
),
|
|
*silence_frames(22, 640, 3),
|
|
]
|
|
transport = PlaybackInjectedTransport(
|
|
first_question_frames,
|
|
inject_after_play_count=18,
|
|
injected_frames=barge_frames,
|
|
)
|
|
stt = QueueSttProvider(["第一问", "打断问题"])
|
|
llm = QueueLlmProvider(
|
|
[
|
|
[
|
|
"已播出一句。",
|
|
"这是一段需要被打断的很长很长很长很长很长很长很长很长的回复内容,没有播放完。",
|
|
],
|
|
["这是新回答。"],
|
|
]
|
|
)
|
|
reporter = RecordingReporter()
|
|
runtime = VoiceAssistantPipeline(
|
|
config=AppConfig(
|
|
llm_api_key="secret",
|
|
speech_provider="cloud",
|
|
wake_ack_text="",
|
|
barge_in_enabled=True,
|
|
barge_in_echo_guard_ms=0,
|
|
barge_in_min_speech_ms=40,
|
|
barge_in_listen_interval_ms=1,
|
|
barge_in_chunk_ms=20,
|
|
followup_listen_timeout_ms=3000,
|
|
),
|
|
transport=transport,
|
|
wakeword=KeywordWakeWordProvider(),
|
|
vad_recorder=VadRecorder(EnergyVadProvider(), min_duration_ms=40, end_silence_ms=40),
|
|
stt=stt,
|
|
realtime_stt=MetadataSttProvider(),
|
|
llm=llm,
|
|
tts=SineTtsProvider(),
|
|
context=ConversationContext(),
|
|
reporter=reporter,
|
|
event_bus=PipelineEventBus(),
|
|
)
|
|
|
|
summary = runtime.run(max_turns=2)
|
|
event_types = [event.type for event in runtime.event_bus.events]
|
|
context_texts = [message.content for message in runtime.context.messages()]
|
|
playback_window = event_types[event_types.index(TTS_STARTED) : event_types.index(PLAYBACK_INTERRUPTED)]
|
|
|
|
self.assertEqual(summary.completed_turns, 2)
|
|
self.assertIn(BARGE_IN_DETECTED, event_types)
|
|
self.assertIn(PLAYBACK_INTERRUPTED, event_types)
|
|
self.assertNotIn(TRANSCRIPT_PARTIAL, playback_window)
|
|
self.assertIn("已播出一句。", context_texts)
|
|
self.assertNotIn("这是一段需要被打断的很长很长很长很长很长很长很长很长的回复内容,没有播放完。", context_texts)
|
|
self.assertEqual(reporter.transcripts[-1], "打断问题")
|
|
self.assertEqual(len(llm.calls), 2)
|
|
|
|
|
|
if __name__ == "__main__":
|
|
unittest.main()
|