import math import numpy as np import pytest from audio_artifact_gate import ( EchoSmearingGateConfig, evaluate_echo_smearing, ) SAMPLE_RATE = 16_000 def _speech_proxy(*, seed=7, seconds=3.0): sample_count = int(SAMPLE_RATE * seconds) time = np.arange(sample_count, dtype=np.float64) / SAMPLE_RATE rng = np.random.default_rng(seed) fundamental = ( 128.0 + 26.0 * np.sin(2.0 * np.pi * 0.37 * time) + 9.0 * np.sin(2.0 * np.pi * 0.11 * time) ) phase = 2.0 * np.pi * np.cumsum(fundamental) / SAMPLE_RATE voiced = sum( np.sin(harmonic * phase + 0.17 * harmonic) / harmonic**0.8 for harmonic in range(1, 18) ) envelope = np.zeros_like(time) for index, center in enumerate(np.linspace(0.20, seconds - 0.25, 8)): width = 0.07 + 0.035 * ((index + seed) % 4) amplitude = 0.60 + 0.20 * ((index + seed) % 3) envelope += amplitude * np.exp(-0.5 * np.square((time - center) / width)) envelope = np.minimum(envelope, 1.0) waveform = voiced * envelope + 0.035 * rng.standard_normal(sample_count) * envelope return 0.60 * waveform / np.max(np.abs(waveform)) def _delayed_copy(waveform, *, delay_ms=73.0, gain=0.68): delay = int(round(delay_ms * SAMPLE_RATE / 1000.0)) output = waveform.copy() output[delay:] += gain * waveform[:-delay] return output / max(1.0, np.max(np.abs(output)) / 0.90) def _dense_smear(waveform): rng = np.random.default_rng(44) impulse = np.zeros(int(0.20 * SAMPLE_RATE), dtype=np.float64) impulse[0] = 1.0 indices = np.arange(int(0.015 * SAMPLE_RATE), impulse.size, 31) impulse[indices] = ( rng.choice((-1.0, 1.0), size=indices.size) * 0.16 * np.exp(-indices / (0.09 * SAMPLE_RATE)) ) output = np.convolve(waveform, impulse, mode="full")[: waveform.size] return output / max(1.0, np.max(np.abs(output)) / 0.90) def _comb_filter(waveform): impulse = np.zeros(int(0.15 * SAMPLE_RATE), dtype=np.float64) impulse[0] = 1.0 for delay_seconds, gain in ( (0.025, 0.50), (0.047, 0.38), (0.071, 0.30), (0.110, 0.20), ): impulse[int(round(delay_seconds * SAMPLE_RATE))] = gain output = np.convolve(waveform, impulse, mode="full")[: waveform.size] return output / max(1.0, np.max(np.abs(output)) / 0.90) @pytest.mark.parametrize("seed", range(5)) def test_clean_nonstationary_speech_proxy_passes(seed): diagnostics = evaluate_echo_smearing(_speech_proxy(seed=seed), SAMPLE_RATE) assert diagnostics.passed assert diagnostics.rejection_reasons == () assert diagnostics.delayed_copy_score < 0.33 assert diagnostics.smearing_score < 0.40 assert diagnostics.analysis_window_count >= 3 def test_clean_broadband_noise_passes(): rng = np.random.default_rng(1234) noise = 0.08 * rng.standard_normal(3 * SAMPLE_RATE) diagnostics = evaluate_echo_smearing(noise, SAMPLE_RATE) assert diagnostics.passed assert diagnostics.rejection_reasons == () @pytest.mark.parametrize("delay_ms", (48.0, 73.0, 121.0)) def test_strong_delayed_copy_is_rejected_and_delay_is_localized(delay_ms): diagnostics = evaluate_echo_smearing( _delayed_copy(_speech_proxy(), delay_ms=delay_ms), SAMPLE_RATE, ) assert not diagnostics.passed assert "delayed_copy" in diagnostics.rejection_reasons assert diagnostics.delayed_copy_score >= 0.45 assert diagnostics.dominant_delay_ms == pytest.approx(delay_ms, abs=0.20) def test_moderate_current_smoke_like_comb_signature_is_rejected(): diagnostics = evaluate_echo_smearing( _delayed_copy(_speech_proxy(), delay_ms=32.0, gain=0.10), SAMPLE_RATE, ) assert not diagnostics.passed assert "delayed_copy" in diagnostics.rejection_reasons assert diagnostics.cepstral_peak_ratio > 10.0 def test_dense_reverberant_smear_is_rejected_by_density_proxy(): diagnostics = evaluate_echo_smearing(_dense_smear(_speech_proxy()), SAMPLE_RATE) assert not diagnostics.passed assert "smearing" in diagnostics.rejection_reasons assert diagnostics.smearing_score >= 0.40 def test_multi_tap_comb_filter_is_rejected(): diagnostics = evaluate_echo_smearing(_comb_filter(_speech_proxy()), SAMPLE_RATE) assert not diagnostics.passed assert "delayed_copy" in diagnostics.rejection_reasons def test_scores_are_deterministic_and_amplitude_invariant(): waveform = _delayed_copy(_speech_proxy()) first = evaluate_echo_smearing(waveform, SAMPLE_RATE) second = evaluate_echo_smearing(waveform.copy(), SAMPLE_RATE) quiet = evaluate_echo_smearing(waveform * 0.125, SAMPLE_RATE) assert first == second assert quiet.delayed_copy_score == pytest.approx( first.delayed_copy_score, abs=1.0e-10 ) assert quiet.smearing_score == pytest.approx(first.smearing_score, abs=1.0e-10) @pytest.mark.parametrize( ("audio", "sample_rate", "reason"), ( ([], SAMPLE_RATE, "invalid_audio"), (np.zeros((2, 8_000)), SAMPLE_RATE, "invalid_audio"), (np.ones(SAMPLE_RATE, dtype=np.complex128), SAMPLE_RATE, "invalid_audio"), (np.full(SAMPLE_RATE, np.nan), SAMPLE_RATE, "nonfinite_audio"), (np.zeros(SAMPLE_RATE), SAMPLE_RATE, "silent_audio"), (np.ones(1_000), SAMPLE_RATE, "insufficient_duration"), (np.ones(SAMPLE_RATE), True, "invalid_sample_rate"), (np.ones(SAMPLE_RATE), 7_999, "invalid_sample_rate"), ), ) def test_invalid_or_insufficient_inputs_fail_closed_with_finite_diagnostics( audio, sample_rate, reason, ): diagnostics = evaluate_echo_smearing(audio, sample_rate) assert not diagnostics.passed assert diagnostics.rejection_reasons == (reason,) for value in ( diagnostics.delayed_copy_score, diagnostics.smearing_score, diagnostics.dominant_delay_ms, diagnostics.cepstral_peak_ratio, diagnostics.cepstral_q99_ratio, diagnostics.cepstral_dense_fraction, diagnostics.duration_seconds, diagnostics.input_rms, ): assert math.isfinite(value) def test_constant_dc_has_no_active_speech_windows_and_fails_closed(): diagnostics = evaluate_echo_smearing(np.full(2 * SAMPLE_RATE, 0.25), SAMPLE_RATE) assert not diagnostics.passed assert diagnostics.rejection_reasons == ("insufficient_active_windows",) def test_invalid_config_fails_closed_instead_of_raising(): diagnostics = evaluate_echo_smearing( _speech_proxy(), SAMPLE_RATE, config=EchoSmearingGateConfig(delayed_copy_threshold=float("nan")), ) assert not diagnostics.passed assert diagnostics.rejection_reasons == ("invalid_config",) malformed = evaluate_echo_smearing( _speech_proxy(), SAMPLE_RATE, config=EchoSmearingGateConfig(min_duration_seconds="not-a-number"), ) assert not malformed.passed assert malformed.rejection_reasons == ("invalid_config",)