[连续对话判断]:完成自动持续对话和播报打断,包含回复意图判断、免唤醒追问和打断回归测试
This commit is contained in:
@@ -0,0 +1,81 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import unittest
|
||||
|
||||
from owner_voice_pet.continuation import (
|
||||
HybridContinuationDecisionProvider,
|
||||
LlmContinuationDecisionProvider,
|
||||
RuleContinuationDecisionProvider,
|
||||
)
|
||||
from owner_voice_pet.models import Message, ReplyDelta
|
||||
|
||||
|
||||
class FakeClassifierLlm:
|
||||
def __init__(self, text: str) -> None:
|
||||
self.text = text
|
||||
self.calls: list[list[Message]] = []
|
||||
|
||||
def stream_reply(self, messages):
|
||||
self.calls.append(list(messages))
|
||||
yield ReplyDelta(self.text, finish_reason="stop")
|
||||
|
||||
|
||||
class ContinuationDecisionTests(unittest.TestCase):
|
||||
def test_rule_continue_for_assistant_question(self) -> None:
|
||||
decision = RuleContinuationDecisionProvider().decide(
|
||||
user_text="讲讲天气",
|
||||
assistant_text="你想继续听哪一部分?",
|
||||
history=[],
|
||||
)
|
||||
|
||||
self.assertEqual(decision.action, "continue")
|
||||
self.assertGreaterEqual(decision.confidence, 0.65)
|
||||
|
||||
def test_rule_standby_for_completed_answer(self) -> None:
|
||||
decision = RuleContinuationDecisionProvider().decide(
|
||||
user_text="今天天气怎么样",
|
||||
assistant_text="这是今天的天气。",
|
||||
history=[],
|
||||
)
|
||||
|
||||
self.assertEqual(decision.action, "standby")
|
||||
self.assertGreaterEqual(decision.confidence, 0.65)
|
||||
|
||||
def test_hybrid_low_confidence_llm_defaults_to_standby(self) -> None:
|
||||
llm = FakeClassifierLlm('{"action":"continue","confidence":0.3,"reason":"弱"}')
|
||||
provider = HybridContinuationDecisionProvider(
|
||||
RuleContinuationDecisionProvider(),
|
||||
LlmContinuationDecisionProvider(llm, threshold=0.65),
|
||||
threshold=0.65,
|
||||
)
|
||||
|
||||
decision = provider.decide(
|
||||
user_text="继续",
|
||||
assistant_text="我还可以从背景原因和下一步影响两个方向继续展开",
|
||||
history=[],
|
||||
)
|
||||
|
||||
self.assertEqual(decision.action, "standby")
|
||||
self.assertEqual(len(llm.calls), 1)
|
||||
|
||||
def test_hybrid_high_confidence_llm_can_continue(self) -> None:
|
||||
llm = FakeClassifierLlm('{"action":"continue","confidence":0.91,"reason":"等待用户选择"}')
|
||||
provider = HybridContinuationDecisionProvider(
|
||||
RuleContinuationDecisionProvider(),
|
||||
LlmContinuationDecisionProvider(llm, threshold=0.65),
|
||||
threshold=0.65,
|
||||
)
|
||||
|
||||
decision = provider.decide(
|
||||
user_text="继续",
|
||||
assistant_text="我还可以从背景原因和下一步影响两个方向继续展开",
|
||||
history=[],
|
||||
)
|
||||
|
||||
self.assertEqual(decision.action, "continue")
|
||||
self.assertEqual(decision.provider, "hybrid")
|
||||
self.assertEqual(len(llm.calls), 1)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
+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()
|
||||
|
||||
@@ -106,6 +106,13 @@ class ModelsConfigTests(unittest.TestCase):
|
||||
self.assertEqual(config.tts_voice, "mimo_default")
|
||||
self.assertEqual(str(config.speech_models_dir), "models")
|
||||
self.assertEqual(config.context_mode, "session_memory")
|
||||
self.assertTrue(config.continuous_dialog_enabled)
|
||||
self.assertEqual(config.continuation_decision_provider, "hybrid")
|
||||
self.assertEqual(config.continuation_confidence_threshold, 0.65)
|
||||
self.assertEqual(config.followup_listen_timeout_ms, 3000)
|
||||
self.assertTrue(config.barge_in_enabled)
|
||||
self.assertEqual(config.barge_in_min_speech_ms, 250)
|
||||
self.assertEqual(config.barge_in_echo_guard_ms, 500)
|
||||
self.assertTrue(config.llm_stream)
|
||||
self.assertEqual(config.validate_basic(), [])
|
||||
|
||||
@@ -139,6 +146,22 @@ class ModelsConfigTests(unittest.TestCase):
|
||||
errors = config.validate_basic()
|
||||
self.assertTrue(any("OWNER_NOISE_FILTER_PROVIDER" in error.message for error in errors))
|
||||
|
||||
def test_continuation_config_is_validated(self) -> None:
|
||||
config = AppConfig(
|
||||
continuation_decision_provider="invalid",
|
||||
continuation_confidence_threshold=2.0,
|
||||
followup_listen_timeout_ms=-1,
|
||||
barge_in_min_speech_ms=-1,
|
||||
barge_in_echo_guard_ms=-1,
|
||||
)
|
||||
errors = config.validate_basic()
|
||||
|
||||
self.assertTrue(any("OWNER_CONTINUATION_DECISION_PROVIDER" in error.message for error in errors))
|
||||
self.assertTrue(any("OWNER_CONTINUATION_CONFIDENCE_THRESHOLD" in error.message for error in errors))
|
||||
self.assertTrue(any("OWNER_FOLLOWUP_LISTEN_TIMEOUT_MS" in error.message for error in errors))
|
||||
self.assertTrue(any("OWNER_BARGE_IN_MIN_SPEECH_MS" in error.message for error in errors))
|
||||
self.assertTrue(any("OWNER_BARGE_IN_ECHO_GUARD_MS" in error.message for error in errors))
|
||||
|
||||
def test_speaker_similarity_threshold_range_is_validated(self) -> None:
|
||||
config = AppConfig(speaker_similarity_threshold=1.5)
|
||||
errors = config.validate_basic()
|
||||
|
||||
Reference in New Issue
Block a user