151 lines
6.2 KiB
Python
151 lines
6.2 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.llm import MockLlmProvider
|
|
from owner_voice_pet.models import AudioFrame, AudioSegment, Transcript
|
|
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]:
|
|
return [
|
|
AudioFrame(b"\xff\x7f", 16000, 1, start_ms, start_id, {"duration_ms": 20, "speech": True}),
|
|
AudioFrame(b"\xff\x7f", 16000, 1, start_ms + 20, start_id + 1, {"duration_ms": 20, "speech": True}),
|
|
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.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}")
|
|
|
|
|
|
def make_runtime(texts: list[str], context: ConversationContext | None = None) -> tuple[LiveVoiceRuntime, 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))
|
|
frames.extend(segment_frames(base_id + 1, base_ms + 20))
|
|
transport = MemoryAudioTransport(frames)
|
|
stt = QueueSttProvider(texts)
|
|
llm = MockLlmProvider(["这是答复。"])
|
|
tts = SineTtsProvider()
|
|
reporter = RecordingReporter()
|
|
runtime = LiveVoiceRuntime(
|
|
config=AppConfig(llm_api_key="secret", speech_provider="cloud", post_playback_drain_ms=0),
|
|
transport=transport,
|
|
wakeword=KeywordWakeWordProvider(),
|
|
vad_recorder=VadRecorder(EnergyVadProvider(), min_duration_ms=40, end_silence_ms=40),
|
|
stt=stt,
|
|
llm=llm,
|
|
tts=tts,
|
|
context=context or ConversationContext(),
|
|
reporter=reporter,
|
|
)
|
|
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(["第一问", "第二问"])
|
|
summary = runtime.run(max_turns=2)
|
|
|
|
self.assertEqual(summary.completed_turns, 2)
|
|
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])
|
|
|
|
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:第一问")
|
|
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()
|