from __future__ import annotations import tempfile import unittest from pathlib import Path from owner_voice_pet.audio_preprocess import SherpaOnnxDenoiserPreprocessor from owner_voice_pet.models import AudioFrame, ErrorCode, ProviderError from owner_voice_pet.speech_models import DEFAULT_DENOISER_PATH class FakeDenoisedAudio: def __init__(self, samples, sample_rate: int) -> None: self.samples = samples self.sample_rate = sample_rate class FakeSherpa: class OfflineSpeechDenoiserGtcrnModelConfig: def __init__(self, model: str) -> None: self.model = model class OfflineSpeechDenoiserModelConfig: def __init__(self, gtcrn, num_threads: int, provider: str) -> None: self.gtcrn = gtcrn self.num_threads = num_threads self.provider = provider class OnlineSpeechDenoiserConfig: def __init__(self, model) -> None: self.model = model class OnlineSpeechDenoiser: def __init__(self, config) -> None: self.config = config self.reset_called = False def run(self, samples, sample_rate: int): import numpy as np return FakeDenoisedAudio(np.zeros_like(samples, dtype=np.float32), sample_rate) def reset(self) -> None: self.reset_called = True class FailingSherpa(FakeSherpa): class OnlineSpeechDenoiser(FakeSherpa.OnlineSpeechDenoiser): def run(self, samples, sample_rate: int): raise RuntimeError("denoise failed") class AudioPreprocessTests(unittest.TestCase): def test_denoiser_missing_model_is_structured(self) -> None: with tempfile.TemporaryDirectory() as tmp: provider = SherpaOnnxDenoiserPreprocessor(tmp, sherpa_module=FakeSherpa) with self.assertRaises(ProviderError) as raised: provider.load() self.assertEqual(raised.exception.code, ErrorCode.NOISE_FILTER_MODEL_MISSING) def test_denoiser_processes_frame_and_marks_metadata(self) -> None: with tempfile.TemporaryDirectory() as tmp: model = Path(tmp) / DEFAULT_DENOISER_PATH model.parent.mkdir(parents=True) model.write_bytes(b"fake") provider = SherpaOnnxDenoiserPreprocessor(tmp, sherpa_module=FakeSherpa) provider.load() frame = AudioFrame( b"\xff\x7f", 16000, 1, 100, 7, {"duration_ms": 20, "speech": True}, ) processed = provider.process_frame(frame) self.assertNotEqual(processed.pcm, frame.pcm) self.assertTrue(processed.metadata["denoised"]) self.assertEqual(processed.metadata["noise_filter_provider"], "sherpa_onnx_gtcrn") self.assertEqual(processed.timestamp_ms, 100) self.assertEqual(processed.frame_id, 7) def test_denoiser_runtime_failure_is_structured(self) -> None: with tempfile.TemporaryDirectory() as tmp: model = Path(tmp) / DEFAULT_DENOISER_PATH model.parent.mkdir(parents=True) model.write_bytes(b"fake") provider = SherpaOnnxDenoiserPreprocessor(tmp, sherpa_module=FailingSherpa) provider.load() with self.assertRaises(ProviderError) as raised: provider.process_frame(AudioFrame(b"\xff\x7f", 16000, 1, 0, 0)) self.assertEqual(raised.exception.code, ErrorCode.NOISE_FILTER_FAILED) if __name__ == "__main__": unittest.main()