343 lines
14 KiB
Python
343 lines
14 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
|
|
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_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 == "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.assertIsNotNone(first)
|
|
self.assertIsNotNone(second)
|
|
assert first is not None and second is not None
|
|
self.assertEqual(first.text, "你")
|
|
self.assertEqual(second.text, "你好")
|
|
|
|
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()
|