[连续对话判断]:完成自动持续对话和播报打断,包含回复意图判断、免唤醒追问和打断回归测试

This commit is contained in:
mkbk
2026-06-18 12:08:41 +08:00
parent 1e956f5eb6
commit 7519725321
16 changed files with 1059 additions and 119 deletions
+192 -1
View File
@@ -6,9 +6,14 @@ 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,
@@ -21,7 +26,7 @@ from owner_voice_pet.events import (
PipelineEventBus,
)
from owner_voice_pet.llm import MockLlmProvider
from owner_voice_pet.models import AudioFrame, AudioSegment, Transcript
from owner_voice_pet.models import AudioFrame, AudioSegment, 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
@@ -95,6 +100,19 @@ class QueueSttProvider:
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 RecordingReporter:
def __init__(self) -> None:
self.statuses: list[str] = []
@@ -188,6 +206,45 @@ def make_runtime(
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)
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_repeated_runtime_runs_two_turns_and_returns_to_standby(self) -> None:
runtime, stt, llm, transport, reporter = make_runtime(["第一问", "第二问"])
@@ -359,6 +416,140 @@ class LiveRuntimeTests(unittest.TestCase):
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)
def test_barge_in_interrupts_playback_and_keeps_only_spoken_context(self) -> None:
first_question_frames = [wake_frame(0, 0)]
first_question_frames.extend(segment_frames(1, 20, partials=["第一问", "第一问"]))
barge_frames = [
AudioFrame(b"\xff\x7f", 16000, 1, 600, 20, {"duration_ms": 20, "speech": True, "partial_transcript": "等一下"}),
AudioFrame(b"\xff\x7f", 16000, 1, 620, 21, {"duration_ms": 20, "speech": True, "partial_transcript": "等一下"}),
*silence_frames(22, 640, 3),
]
transport = PlaybackInjectedTransport(
first_question_frames,
inject_after_play_count=8,
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=500,
barge_in_min_speech_ms=40,
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()]
self.assertEqual(summary.completed_turns, 2)
self.assertIn(BARGE_IN_DETECTED, event_types)
self.assertIn(PLAYBACK_INTERRUPTED, event_types)
self.assertIn("已播出一句。", context_texts)
self.assertNotIn("这是一段需要被打断的很长很长很长很长很长很长很长很长的回复内容,没有播放完。", context_texts)
self.assertEqual(reporter.transcripts[-1], "打断问题")
self.assertEqual(len(llm.calls), 2)
if __name__ == "__main__":
unittest.main()