[Wake/VAD/STT 与 Live runtime]:完成真实重复语音运行,包含云端ASR/TTS开关、run-live和临时上下文测试

This commit is contained in:
mkbk
2026-06-17 20:00:55 +08:00
parent ac97daa1e7
commit 4b7cd18a0f
20 changed files with 1043 additions and 68 deletions
+20 -1
View File
@@ -57,11 +57,18 @@ class CliAcceptanceTests(unittest.TestCase):
path = root / relative
path.parent.mkdir(parents=True, exist_ok=True)
path.write_bytes(b"placeholder")
with patch("importlib.util.find_spec", return_value=object()):
with (
patch("importlib.util.find_spec", return_value=object()),
patch("owner_voice_pet.cli.SherpaOnnxVadProvider") as vad_cls,
patch("owner_voice_pet.cli.SherpaOnnxSttProvider") as stt_cls,
):
vad_cls.return_value.load.return_value = None
stt_cls.return_value.load.return_value = None
code, data = self.call("model-check", "--models-dir", str(root))
self.assertEqual(code, 0)
self.assertTrue(data["ok"])
self.assertEqual(data["missing_files"], [])
self.assertTrue(data["provider_load_checked"])
def test_model_check_reports_missing_files(self) -> None:
with tempfile.TemporaryDirectory() as tmp:
@@ -77,6 +84,18 @@ class CliAcceptanceTests(unittest.TestCase):
self.assertEqual(code, 0)
self.assertTrue(data["ok"])
def test_run_live_once_invokes_runtime(self) -> None:
class FakeRuntime:
def run(self, *, once: bool = False):
self.once = once
return type("Summary", (), {"completed_turns": 1, "interrupted": False})()
fake_runtime = FakeRuntime()
with patch("owner_voice_pet.cli.build_live_runtime", return_value=fake_runtime):
code = main(["run-live", "--once"])
self.assertEqual(code, 0)
self.assertTrue(fake_runtime.once)
if __name__ == "__main__":
unittest.main()
+107
View File
@@ -0,0 +1,107 @@
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
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}),
]
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.errors: list[str] = []
def status(self, state: str, message: str, *, turn_id: int | None = None) -> None:
self.statuses.append(message)
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 in range(4):
frames.extend(segment_frames(idx * 4, idx * 80))
transport = MemoryAudioTransport(frames)
stt = QueueSttProvider(texts)
llm = MockLlmProvider(["这是答复。"])
tts = SineTtsProvider()
reporter = RecordingReporter()
runtime = LiveVoiceRuntime(
config=AppConfig(llm_api_key="secret", speech_provider="cloud"),
transport=transport,
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), 4)
self.assertEqual(len(llm.calls), 2)
self.assertEqual(len(transport.played_segments), 2)
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(), ())
if __name__ == "__main__":
unittest.main()
+24 -3
View File
@@ -59,25 +59,46 @@ class ModelsConfigTests(unittest.TestCase):
handle.write(
"\n".join(
[
"OWNER_LLM_BASE_URL=https://newapi.mkbk.shop/",
"OWNER_LLM_BASE_URL=https://token-plan-cn.xiaomimimo.com/v1/",
"OWNER_LLM_API_KEY=secret-value",
"OWNER_LLM_MODEL=test-model",
]
)
)
config = AppConfig.from_dotenv(path)
self.assertEqual(config.llm_base_url, "https://newapi.mkbk.shop")
self.assertEqual(config.llm_base_url, "https://token-plan-cn.xiaomimimo.com/v1")
self.assertEqual(config.llm_api_key, "secret-value")
self.assertEqual(config.llm_model, "test-model")
self.assertEqual(config.speech_provider, "cloud")
self.assertEqual(config.asr_model, "mimo-v2.5-asr")
self.assertEqual(config.tts_model, "mimo-v2.5-tts")
self.assertEqual(config.tts_voice, "alloy")
self.assertEqual(str(config.speech_models_dir), "models")
self.assertTrue(config.llm_stream)
self.assertEqual(config.validate_basic(), [])
def test_speech_provider_must_be_cloud_or_local(self) -> None:
config = AppConfig(speech_provider="invalid")
errors = config.validate_basic()
self.assertTrue(any("OWNER_SPEECH_PROVIDER" in error.message for error in errors))
def test_missing_dotenv_uses_non_secret_defaults(self) -> None:
config = AppConfig.from_dotenv("/tmp/owner-voice-pet-missing.env")
self.assertEqual(config.llm_base_url, "https://newapi.mkbk.shop")
self.assertEqual(config.llm_base_url, "https://token-plan-cn.xiaomimimo.com/v1")
self.assertIsNone(config.llm_api_key)
def test_api_url_accepts_base_with_or_without_v1(self) -> None:
with_v1 = AppConfig(llm_base_url="https://token-plan-cn.xiaomimimo.com/v1")
without_v1 = AppConfig(llm_base_url="https://newapi.mkbk.shop")
self.assertEqual(
with_v1.api_url("/v1/chat/completions"),
"https://token-plan-cn.xiaomimimo.com/v1/chat/completions",
)
self.assertEqual(
without_v1.api_url("/v1/chat/completions"),
"https://newapi.mkbk.shop/v1/chat/completions",
)
def test_missing_llm_key_has_structured_error(self) -> None:
config = AppConfig(llm_api_key=None)
with self.assertRaises(ProviderError) as raised:
+33 -3
View File
@@ -11,7 +11,7 @@ from owner_voice_pet.models import AudioFrame, ErrorCode, Message, PipelineState
from owner_voice_pet.pipeline import VoicePipeline
from owner_voice_pet.stt import MetadataSttProvider
from owner_voice_pet.transport import MemoryAudioTransport
from owner_voice_pet.tts import SentenceBuffer, SineTtsProvider
from owner_voice_pet.tts import CloudTtsProvider, SentenceBuffer, SineTtsProvider
from owner_voice_pet.vad import EnergyVadProvider, VadRecorder
from owner_voice_pet.wakeword import KeywordWakeWordProvider
@@ -66,6 +66,36 @@ class PipelineLlmTtsTests(unittest.TestCase):
self.assertGreater(len(segment.pcm), 0)
self.assertGreater(segment.duration_ms, 0)
def test_cloud_tts_posts_speech_request(self) -> None:
class FakeResponse:
def __enter__(self):
return self
def __exit__(self, *args) -> None:
return None
def read(self) -> bytes:
return b"fake-mp3"
requests = []
def fake_urlopen(request, timeout):
requests.append(request)
return FakeResponse()
provider = CloudTtsProvider(
AppConfig(llm_api_key="secret", tts_model="mimo-v2.5-tts"),
urlopen=fake_urlopen,
)
provider.load()
segment = provider.synthesize("你好")
self.assertEqual(segment.metadata["format"], "mp3")
self.assertEqual(segment.pcm, b"fake-mp3")
body = json.loads(requests[0].data.decode())
self.assertEqual(body["model"], "mimo-v2.5-tts")
self.assertIn("/v1/audio/speech", requests[0].full_url)
def test_pipeline_runs_from_wake_to_playback(self) -> None:
frames = [
frame(0, 0, {"wake_word": "小杰小杰", "wake_confidence": 0.95}),
@@ -136,7 +166,7 @@ class PipelineLlmTtsTests(unittest.TestCase):
return FakeResponse()
config = AppConfig(
llm_base_url="https://newapi.mkbk.shop",
llm_base_url="https://token-plan-cn.xiaomimimo.com/v1",
llm_api_key="secret",
llm_model="test-model",
llm_api_style="chat_completions",
@@ -160,7 +190,7 @@ class PipelineLlmTtsTests(unittest.TestCase):
return json.dumps({"choices": [{"message": {"content": "非流式回复。"}}]}).encode()
config = AppConfig(
llm_base_url="https://newapi.mkbk.shop",
llm_base_url="https://token-plan-cn.xiaomimimo.com/v1",
llm_api_key="secret",
llm_model="test-model",
llm_api_style="chat_completions",
+31 -1
View File
@@ -1,10 +1,12 @@
from __future__ import annotations
import json
import tempfile
import unittest
from owner_voice_pet.models import AudioFrame, AudioSegment, ErrorCode, ProviderError
from owner_voice_pet.stt import MetadataSttProvider, SherpaOnnxSttProvider, is_valid_transcript_text
from owner_voice_pet.config import AppConfig
from owner_voice_pet.stt import CloudAsrSttProvider, MetadataSttProvider, SherpaOnnxSttProvider, is_valid_transcript_text
from owner_voice_pet.vad import EnergyVadProvider, VadRecorder
from owner_voice_pet.wakeword import KeywordWakeWordProvider, MissingWakeWordModelProvider
@@ -102,6 +104,34 @@ class WakeVadSttTests(unittest.TestCase):
provider.load()
self.assertEqual(raised.exception.code, ErrorCode.STT_MODEL_MISSING)
def test_cloud_asr_posts_audio_transcription_request(self) -> None:
class FakeResponse:
def __enter__(self):
return self
def __exit__(self, *args) -> None:
return None
def read(self) -> bytes:
return json.dumps({"text": "你好小杰", "language": "zh"}, ensure_ascii=False).encode()
requests = []
def fake_urlopen(request, timeout):
requests.append(request)
return FakeResponse()
provider = CloudAsrSttProvider(
AppConfig(llm_api_key="secret", asr_model="mimo-v2.5-asr"),
urlopen=fake_urlopen,
)
provider.load()
transcript = provider.transcribe(AudioSegment(b"\x00\x00\x01\x00", 16000, 1, 0, 100))
self.assertEqual(transcript.text, "你好小杰")
self.assertIn("/v1/audio/transcriptions", requests[0].full_url)
self.assertIn(b'mimo-v2.5-asr', requests[0].data)
if __name__ == "__main__":
unittest.main()