from __future__ import annotations import unittest from unittest.mock import patch from owner_voice_pet.config import AppConfig from owner_voice_pet.full_duplex_audio import ( CaptureRingBuffer, FakeWebRtcAudioProcessingProvider, NoopAudioProcessingProvider, RenderReferenceRingBuffer, build_audio_processing_provider, webrtc_apm_probe, ) from owner_voice_pet.models import AudioFrame, ErrorCode, ProviderError def frame( frame_id: int, timestamp_ms: int, *, duration_ms: int = 20, pcm: bytes = b"\x01\x00" * 160, sample_rate: int = 16000, channels: int = 1, metadata: dict[str, object] | None = None, ) -> AudioFrame: data = {"duration_ms": duration_ms} if metadata: data.update(metadata) return AudioFrame( pcm=pcm, sample_rate=sample_rate, channels=channels, timestamp_ms=timestamp_ms, frame_id=frame_id, metadata=data, ) class FullDuplexAudioTests(unittest.TestCase): def test_audio_frame_duration_and_fixture_roundtrip(self) -> None: original = frame(7, 140, metadata={"speech": True}) self.assertEqual(original.duration_ms, 20) restored = AudioFrame.from_fixture(original.to_fixture()) self.assertEqual(restored, original) self.assertTrue(restored.metadata["speech"]) def test_capture_ring_buffer_drops_oldest_frames_on_overrun(self) -> None: buffer = CaptureRingBuffer(capacity_ms=40) first = buffer.write(frame(1, 0)) second = buffer.write(frame(2, 20)) third = buffer.write(frame(3, 40)) self.assertFalse(first.overrun) self.assertFalse(second.overrun) self.assertTrue(third.overrun) self.assertEqual([item.frame_id for item in third.dropped], [1]) self.assertEqual([item.frame_id for item in buffer.frames()], [2, 3]) self.assertEqual(buffer.duration_ms, 40) def test_render_reference_window_uses_timestamps(self) -> None: buffer = RenderReferenceRingBuffer(capacity_ms=100) for idx, timestamp in enumerate([0, 20, 40, 60], start=1): buffer.write(frame(idx, timestamp)) window = buffer.latest_window(timestamp_ms=60, window_ms=30) self.assertEqual([item.frame_id for item in window], [3, 4]) def test_fake_webrtc_apm_suppresses_assistant_echo(self) -> None: provider = FakeWebRtcAudioProcessingProvider(sample_rate=16000, channels=1) provider.process_render(frame(1, 0, metadata={"assistant_audio": True})) processed = provider.process_capture(frame(2, 20, metadata={"assistant_echo": True, "speech": True})) self.assertEqual(processed.pcm, b"\x00" * len(processed.pcm)) self.assertTrue(processed.metadata["echo_suppressed"]) self.assertFalse(processed.metadata["speech"]) def test_fake_webrtc_apm_rejects_format_mismatch(self) -> None: provider = FakeWebRtcAudioProcessingProvider(sample_rate=16000, channels=1) with self.assertRaises(ProviderError) as raised: provider.process_capture(frame(1, 0, sample_rate=8000)) self.assertEqual(raised.exception.code, ErrorCode.AUDIO_APM_FORMAT_MISMATCH) def test_fake_webrtc_apm_can_report_processing_failure(self) -> None: provider = FakeWebRtcAudioProcessingProvider(fail_processing=True) with self.assertRaises(ProviderError) as raised: provider.process_capture(frame(1, 0)) self.assertEqual(raised.exception.code, ErrorCode.AUDIO_APM_PROCESS_FAILED) def test_audio_processing_provider_uses_fake_provider(self) -> None: provider = build_audio_processing_provider(AppConfig(audio_apm_provider="fake")) self.assertIsInstance(provider, FakeWebRtcAudioProcessingProvider) self.assertTrue(provider.health_check().available) def test_required_webrtc_provider_fails_when_probe_is_unavailable(self) -> None: with patch("owner_voice_pet.full_duplex_audio.webrtc_apm_probe") as probe: probe.return_value = type( "Health", (), {"available": False, "message": "missing binding"}, )() with self.assertRaises(ProviderError) as raised: build_audio_processing_provider( AppConfig(audio_apm_provider="webrtc", audio_apm_required=True) ) self.assertEqual(raised.exception.code, ErrorCode.AUDIO_APM_UNAVAILABLE) def test_webrtc_provider_can_fallback_when_not_required(self) -> None: with patch("owner_voice_pet.full_duplex_audio.webrtc_apm_probe") as probe: probe.return_value = type( "Health", (), {"available": False, "message": "missing binding"}, )() provider = build_audio_processing_provider( AppConfig(audio_apm_provider="webrtc", audio_apm_required=False) ) self.assertIsInstance(provider, NoopAudioProcessingProvider) self.assertTrue(provider.health_check().fallback_active) def test_webrtc_probe_reports_unavailable_without_binding(self) -> None: with patch("importlib.util.find_spec", return_value=None): health = webrtc_apm_probe() self.assertFalse(health.available) self.assertEqual(health.provider, "webrtc") if __name__ == "__main__": unittest.main()