Files
Owner/tests/test_live_runtime.py
T

278 lines
11 KiB
Python

from __future__ import annotations
import unittest
from owner_voice_pet.config import AppConfig
from owner_voice_pet.conversation import ConversationContext
from owner_voice_pet.events import (
ACK_STARTED,
CAPTURE_STARTED,
LLM_STARTED,
PLAYBACK_FINISHED,
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, Transcript
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 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 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,
) -> 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)
stt = QueueSttProvider(texts)
llm = MockLlmProvider(["这是答复。"])
tts = SineTtsProvider()
reporter = RecordingReporter()
event_bus = PipelineEventBus()
runtime = VoiceAssistantPipeline(
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,
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 LiveRuntimeTests(unittest.TestCase):
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), 4)
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_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_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_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)
if __name__ == "__main__":
unittest.main()