[连续对话判断]:完成自动持续对话和播报打断,包含回复意图判断、免唤醒追问和打断回归测试
This commit is contained in:
+192
-1
@@ -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()
|
||||
|
||||
Reference in New Issue
Block a user