Files
Owner/tests/test_wake_vad_stt.py

404 lines
16 KiB
Python

from __future__ import annotations
import json
import tempfile
import unittest
from pathlib import Path
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,
should_emit_partial_transcript,
)
from owner_voice_pet.speech_models import stt_model_paths
from owner_voice_pet.vad import EnergyVadProvider, HybridVadProvider, PrimarySpeakerVadRecorder, 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_hybrid_vad_accepts_energy_fallback_speech(self) -> None:
class SilentVadProvider:
def __init__(self) -> None:
self.loaded = False
def load(self) -> None:
self.loaded = True
def analyze(self, frame: AudioFrame):
return type(
"Result",
(),
{"is_speech": False, "confidence": 0.1, "speech_ms": 0, "silence_ms": 20},
)()
def reset(self) -> None:
return None
provider = HybridVadProvider(SilentVadProvider(), EnergyVadProvider())
provider.load()
result = provider.analyze(make_frame(1, 0, speech=True))
self.assertTrue(result.is_speech)
self.assertEqual(result.speech_ms, 20)
def test_primary_speaker_endpoint_stops_on_background_noise(self) -> None:
provider = EnergyVadProvider()
provider.load()
recorder = PrimarySpeakerVadRecorder(
provider,
min_duration_ms=250,
end_silence_ms=1000,
speaker_profile_ms=40,
speaker_profile_min_ms=40,
speaker_absent_ms=40,
)
frames = [
make_frame(1, 0, speech=True, metadata={"speaker_id": "owner", "transcript": "你是谁"}),
make_frame(2, 20, speech=True, metadata={"speaker_id": "owner"}),
make_frame(3, 40, speech=True, metadata={"speaker_id": "background"}),
make_frame(4, 60, speech=True, metadata={"speaker_id": "background"}),
make_frame(5, 80, speech=True, metadata={"speaker_id": "owner", "transcript": "你是谁第二次"}),
]
segment = None
consumed = 0
for item in frames:
consumed += 1
result = recorder.feed(item)
if isinstance(result, AudioSegment):
segment = result
break
self.assertIsNotNone(segment)
assert segment is not None
self.assertEqual(segment.metadata["end_reason"], "primary_speaker_absent")
self.assertEqual(consumed, 4)
self.assertEqual(segment.metadata["transcript"], "你是谁")
def test_primary_speaker_endpoint_does_not_wait_for_vad_min_duration(self) -> None:
provider = EnergyVadProvider()
provider.load()
recorder = PrimarySpeakerVadRecorder(
provider,
min_duration_ms=1000,
end_silence_ms=1000,
speaker_profile_ms=120,
speaker_profile_min_ms=40,
speaker_absent_ms=40,
)
frames = [
make_frame(1, 0, speech=True, metadata={"speaker_id": "owner", "transcript": "你在做什么"}),
make_frame(2, 20, speech=True, metadata={"speaker_id": "owner"}),
make_frame(3, 40, speech=True, metadata={"speaker_id": "background"}),
make_frame(4, 60, speech=True, metadata={"speaker_id": "background"}),
make_frame(5, 80, speech=True, metadata={"speaker_id": "owner", "transcript": "第二次重复"}),
]
segment = None
for item in frames:
result = recorder.feed(item)
if isinstance(result, AudioSegment):
segment = result
break
self.assertIsNotNone(segment)
assert segment is not None
self.assertEqual(segment.metadata["end_reason"], "primary_speaker_absent")
self.assertEqual(segment.metadata["transcript"], "你在做什么")
def test_primary_speaker_endpoint_allows_short_pause(self) -> None:
provider = EnergyVadProvider()
provider.load()
recorder = PrimarySpeakerVadRecorder(
provider,
min_duration_ms=40,
end_silence_ms=1000,
speaker_profile_ms=40,
speaker_profile_min_ms=40,
speaker_absent_ms=60,
)
frames = [
make_frame(1, 0, speech=True, metadata={"speaker_id": "owner"}),
make_frame(2, 20, speech=True, metadata={"speaker_id": "owner"}),
make_frame(3, 40, speech=False),
make_frame(4, 60, speech=True, metadata={"speaker_id": "owner"}),
make_frame(5, 80, speech=False),
make_frame(6, 100, speech=False),
make_frame(7, 120, speech=False),
]
results = [recorder.feed(item) for item in frames]
self.assertIsNone(results[2])
self.assertIsNone(results[3])
self.assertIsInstance(results[-1], AudioSegment)
assert isinstance(results[-1], AudioSegment)
self.assertEqual(results[-1].metadata["end_reason"], "primary_speaker_absent")
def test_primary_speaker_endpoint_falls_back_to_vad_when_profile_missing(self) -> None:
provider = EnergyVadProvider()
provider.load()
recorder = PrimarySpeakerVadRecorder(provider, min_duration_ms=40, end_silence_ms=40)
frames = [
AudioFrame(b"", 16000, 1, 0, 1, {"duration_ms": 20, "speech": True, "transcript": "你好"}),
AudioFrame(b"", 16000, 1, 20, 2, {"duration_ms": 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)
assert segment is not None
self.assertEqual(segment.metadata["end_reason"], "silence")
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_partial_transcript_filter_rejects_short_noise(self) -> None:
self.assertFalse(should_emit_partial_transcript("", ""))
self.assertFalse(should_emit_partial_transcript("家确", ""))
self.assertTrue(should_emit_partial_transcript("你是谁", ""))
self.assertFalse(should_emit_partial_transcript("", "你是谁"))
self.assertTrue(should_emit_partial_transcript("你是谁呀", "你是谁"))
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_sherpa_stt_streaming_session_emits_partial_text(self) -> None:
class FakeResult:
def __init__(self, text: str) -> None:
self.text = text
class FakeStream:
def __init__(self) -> None:
self.ready = False
self.text = ""
def accept_waveform(self, sample_rate, samples) -> None:
self.ready = True
self.text = "" if not self.text else "你好吗"
class FakeRecognizer:
def create_stream(self):
return FakeStream()
def is_ready(self, stream) -> bool:
return stream.ready
def decode_stream(self, stream) -> None:
stream.ready = False
def get_result(self, stream):
return FakeResult(stream.text)
class FakeOnlineRecognizer:
@staticmethod
def from_transducer(**kwargs):
return FakeRecognizer()
class FakeSherpa:
OnlineRecognizer = FakeOnlineRecognizer
with tempfile.TemporaryDirectory() as tmp:
paths = stt_model_paths(Path(tmp))
for name, path in paths.items():
if name in {"type", "model_dir"}:
continue
path.parent.mkdir(parents=True, exist_ok=True)
path.write_text("fake", encoding="utf-8")
provider = SherpaOnnxSttProvider(tmp, sherpa_module=FakeSherpa)
provider.load()
session = provider.start_stream()
first = session.accept_frame(make_frame(1, 0, speech=True))
second = session.accept_frame(make_frame(2, 20, speech=True))
self.assertIsNone(first)
self.assertIsNotNone(second)
assert second is not None
self.assertEqual(second.text, "你好吗")
def test_sherpa_stt_loads_ctc_manifest(self) -> None:
class FakeRecognizer:
def create_stream(self):
return object()
class FakeOnlineRecognizer:
ctc_kwargs = None
@staticmethod
def from_zipformer2_ctc(**kwargs):
FakeOnlineRecognizer.ctc_kwargs = kwargs
return FakeRecognizer()
class FakeSherpa:
OnlineRecognizer = FakeOnlineRecognizer
with tempfile.TemporaryDirectory() as tmp:
root = Path(tmp)
stt_dir = root / "stt" / "ctc"
stt_dir.mkdir(parents=True)
(stt_dir / "tokens.txt").write_text("你 1\n", encoding="utf-8")
(stt_dir / "model.int8.onnx").write_bytes(b"fake")
(root / "manifest.json").write_text(
json.dumps(
{
"providers": {
"stt": {
"type": "sherpa-onnx-streaming-zipformer2-ctc",
"model_dir": "stt/ctc",
"tokens": "stt/ctc/tokens.txt",
"model": "stt/ctc/model.int8.onnx",
}
},
"required_files": [
"stt/ctc/tokens.txt",
"stt/ctc/model.int8.onnx",
],
},
ensure_ascii=False,
),
encoding="utf-8",
)
provider = SherpaOnnxSttProvider(tmp, sherpa_module=FakeSherpa)
provider.load()
self.assertIsNotNone(FakeOnlineRecognizer.ctc_kwargs)
assert FakeOnlineRecognizer.ctc_kwargs is not None
self.assertTrue(FakeOnlineRecognizer.ctc_kwargs["model"].endswith("model.int8.onnx"))
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()