156 lines
6.2 KiB
Python
156 lines
6.2 KiB
Python
from __future__ import annotations
|
|
|
|
import json
|
|
import tempfile
|
|
import unittest
|
|
|
|
from owner_voice_pet.models import AudioFrame, AudioSegment, ErrorCode, ProviderError
|
|
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,
|
|
SherpaOnnxKeywordWakeWordProvider,
|
|
)
|
|
|
|
|
|
def make_frame(
|
|
idx: int,
|
|
timestamp_ms: int,
|
|
*,
|
|
speech: bool = False,
|
|
metadata: dict[str, object] | None = None,
|
|
) -> AudioFrame:
|
|
data = b"\xff\xff" if speech else b"\x80\x80"
|
|
merged = {"speech": speech, "duration_ms": 20}
|
|
if metadata:
|
|
merged.update(metadata)
|
|
return AudioFrame(data, 16000, 1, timestamp_ms, idx, merged)
|
|
|
|
|
|
class WakeVadSttTests(unittest.TestCase):
|
|
def test_keyword_wakeword_detects_chinese_phrase(self) -> None:
|
|
provider = KeywordWakeWordProvider("小杰小杰", threshold=0.7)
|
|
provider.load()
|
|
event = provider.detect(
|
|
make_frame(1, 100, metadata={"wake_word": "小杰小杰", "wake_confidence": 0.9})
|
|
)
|
|
self.assertIsNotNone(event)
|
|
self.assertEqual(event.keyword, "小杰小杰")
|
|
|
|
def test_keyword_wakeword_ignores_low_confidence(self) -> None:
|
|
provider = KeywordWakeWordProvider("小杰小杰", threshold=0.8)
|
|
provider.load()
|
|
self.assertIsNone(
|
|
provider.detect(make_frame(1, 100, metadata={"wake": True, "wake_confidence": 0.2}))
|
|
)
|
|
|
|
def test_missing_wake_model_reports_structured_error(self) -> None:
|
|
with self.assertRaises(ProviderError) as raised:
|
|
MissingWakeWordModelProvider("/missing/model.onnx").load()
|
|
self.assertEqual(raised.exception.code, ErrorCode.WAKE_MODEL_MISSING)
|
|
|
|
def test_sherpa_kws_missing_model_is_structured(self) -> None:
|
|
with tempfile.TemporaryDirectory() as tmp:
|
|
provider = SherpaOnnxKeywordWakeWordProvider(tmp)
|
|
with self.assertRaises(ProviderError) as raised:
|
|
provider.load()
|
|
self.assertEqual(raised.exception.code, ErrorCode.WAKE_MODEL_MISSING)
|
|
|
|
def test_vad_recorder_returns_segment_after_silence(self) -> None:
|
|
provider = EnergyVadProvider()
|
|
provider.load()
|
|
recorder = VadRecorder(provider, min_duration_ms=40, end_silence_ms=40)
|
|
frames = [
|
|
make_frame(1, 0, speech=True, metadata={"transcript": "你好"}),
|
|
make_frame(2, 20, speech=True),
|
|
make_frame(3, 40, speech=False),
|
|
make_frame(4, 60, speech=False),
|
|
]
|
|
segment = None
|
|
for item in frames:
|
|
result = recorder.feed(item)
|
|
if isinstance(result, AudioSegment):
|
|
segment = result
|
|
self.assertIsNotNone(segment)
|
|
self.assertEqual(segment.metadata["end_reason"], "silence")
|
|
self.assertEqual(segment.metadata["transcript"], "你好")
|
|
|
|
def test_vad_recorder_returns_no_speech_timeout_error(self) -> None:
|
|
provider = EnergyVadProvider()
|
|
provider.load()
|
|
recorder = VadRecorder(provider, no_speech_timeout_ms=40)
|
|
result = None
|
|
for item in [make_frame(1, 0), make_frame(2, 20), make_frame(3, 40)]:
|
|
result = recorder.feed(item)
|
|
self.assertIsInstance(result, ProviderError)
|
|
self.assertEqual(result.code, ErrorCode.VAD_TIMEOUT_NO_SPEECH)
|
|
|
|
def test_metadata_stt_transcribes_fixture_text(self) -> None:
|
|
provider = MetadataSttProvider()
|
|
provider.load()
|
|
transcript = provider.transcribe(
|
|
AudioSegment(b"\x01\x00", 16000, 1, 0, 500, {"transcript": "今天天气怎么样"})
|
|
)
|
|
self.assertEqual(transcript.normalized_text, "今天天气怎么样")
|
|
self.assertEqual(transcript.language, "zh")
|
|
|
|
def test_metadata_stt_rejects_empty_text(self) -> None:
|
|
provider = MetadataSttProvider()
|
|
provider.load()
|
|
with self.assertRaises(ProviderError) as raised:
|
|
provider.transcribe(AudioSegment(b"\x01\x00", 16000, 1, 0, 500, {"transcript": "。!?"}))
|
|
self.assertEqual(raised.exception.code, ErrorCode.STT_EMPTY_TRANSCRIPT)
|
|
|
|
def test_transcript_validator_accepts_chinese_and_ascii(self) -> None:
|
|
self.assertTrue(is_valid_transcript_text("你好"))
|
|
self.assertTrue(is_valid_transcript_text("hello"))
|
|
self.assertFalse(is_valid_transcript_text("?! 。"))
|
|
|
|
def test_sherpa_stt_missing_model_is_structured(self) -> None:
|
|
with tempfile.TemporaryDirectory() as tmp:
|
|
provider = SherpaOnnxSttProvider(f"{tmp}/missing")
|
|
with self.assertRaises(ProviderError) as raised:
|
|
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(
|
|
{"choices": [{"message": {"content": "你好小杰"}}]},
|
|
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/chat/completions", requests[0].full_url)
|
|
body = json.loads(requests[0].data.decode())
|
|
self.assertEqual(body["model"], "mimo-v2.5-asr")
|
|
content = body["messages"][0]["content"][0]
|
|
self.assertEqual(content["type"], "input_audio")
|
|
self.assertTrue(content["input_audio"]["data"].startswith("data:audio/wav;base64,"))
|
|
|
|
|
|
if __name__ == "__main__":
|
|
unittest.main()
|