[Wake/VAD/STT 与 Live runtime]:完成真实重复语音运行,包含云端ASR/TTS开关、run-live和临时上下文测试
This commit is contained in:
@@ -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()
|
||||
|
||||
@@ -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()
|
||||
@@ -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:
|
||||
|
||||
@@ -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",
|
||||
|
||||
@@ -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()
|
||||
|
||||
Reference in New Issue
Block a user