[独立唤醒与转写显示]:完成运行时本地唤醒分离,包含wake注入、转写输出和污染回归测试

This commit is contained in:
mkbk
2026-06-17 20:36:10 +08:00
parent e565164e6e
commit 4b21e0c346
3 changed files with 104 additions and 68 deletions
+48 -9
View File
@@ -10,6 +10,7 @@ from owner_voice_pet.runtime import LiveVoiceRuntime
from owner_voice_pet.transport import MemoryAudioTransport
from owner_voice_pet.tts import SineTtsProvider
from owner_voice_pet.vad import EnergyVadProvider, VadRecorder
from owner_voice_pet.wakeword import KeywordWakeWordProvider
def segment_frames(start_id: int, start_ms: int) -> list[AudioFrame]:
@@ -21,6 +22,17 @@ def segment_frames(start_id: int, start_ms: int) -> list[AudioFrame]:
]
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)
@@ -39,10 +51,17 @@ class QueueSttProvider:
class RecordingReporter:
def __init__(self) -> None:
self.statuses: list[str] = []
self.transcripts: 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:
self.transcripts.append(text)
self.events.append(f"transcript:{text}")
def error(self, stage: str, code: str, message: str, *, turn_id: int | None = None) -> None:
self.errors.append(f"{stage}:{code}:{message}")
@@ -50,8 +69,11 @@ class RecordingReporter:
def make_runtime(texts: list[str], context: ConversationContext | None = None) -> tuple[LiveVoiceRuntime, QueueSttProvider, MockLlmProvider, MemoryAudioTransport, RecordingReporter]:
frames = []
for idx in range(4):
frames.extend(segment_frames(idx * 4, idx * 80))
for idx, _text in enumerate(texts):
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)
stt = QueueSttProvider(texts)
llm = MockLlmProvider(["这是答复。"])
@@ -60,6 +82,7 @@ def make_runtime(texts: list[str], context: ConversationContext | None = None) -
runtime = LiveVoiceRuntime(
config=AppConfig(llm_api_key="secret", speech_provider="cloud"),
transport=transport,
wakeword=KeywordWakeWordProvider(),
vad_recorder=VadRecorder(EnergyVadProvider(), min_duration_ms=40, end_silence_ms=40),
stt=stt,
llm=llm,
@@ -72,19 +95,18 @@ def make_runtime(texts: list[str], context: ConversationContext | None = None) -
class LiveRuntimeTests(unittest.TestCase):
def test_repeated_runtime_runs_two_turns_and_returns_to_standby(self) -> None:
runtime, stt, llm, transport, reporter = make_runtime(
["小杰小杰", "第一问", "小杰小杰", "第二问"]
)
runtime, stt, llm, transport, reporter = make_runtime(["第一问", "第二问"])
summary = runtime.run(max_turns=2)
self.assertEqual(summary.completed_turns, 2)
self.assertEqual(len(stt.calls), 4)
self.assertEqual(len(stt.calls), 2)
self.assertEqual(len(llm.calls), 2)
self.assertEqual(len(transport.played_segments), 2)
self.assertEqual(reporter.transcripts, ["第一问", "第二问"])
self.assertIn("恢复待机:可继续唤醒", reporter.statuses[-1])
def test_temporary_context_is_sent_to_second_llm_call(self) -> None:
runtime, _, llm, _, _ = make_runtime(["小杰小杰", "第一问", "小杰小杰", "第二问"])
runtime, _, llm, _, _ = make_runtime(["第一问", "第二问"])
runtime.run(max_turns=2)
second_call_text = [message.content for message in llm.calls[1]]
@@ -94,14 +116,31 @@ class LiveRuntimeTests(unittest.TestCase):
def test_new_runtime_context_starts_empty(self) -> None:
first_context = ConversationContext()
first_runtime, _, _, _, _ = make_runtime(["小杰小杰", "第一问"], context=first_context)
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)
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:第一问")
thinking_index = next(
index for index, event in enumerate(reporter.events) if event == "status:思考中:正在生成回复"
)
self.assertLess(transcript_index, thinking_index)
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)
if __name__ == "__main__":
unittest.main()