208 lines
8.3 KiB
Python
208 lines
8.3 KiB
Python
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 (
|
|
AudioHub,
|
|
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_audio_hub_fans_out_processed_capture_without_stealing_frames(self) -> None:
|
|
hub = AudioHub(
|
|
processor=FakeWebRtcAudioProcessingProvider(sample_rate=16000, channels=1),
|
|
capture_capacity_ms=100,
|
|
)
|
|
vad = hub.subscribe("processed_capture", name="vad")
|
|
stt = hub.subscribe("processed_capture", name="stt")
|
|
|
|
for idx, timestamp in enumerate([0, 20, 40], start=1):
|
|
hub.accept_capture(frame(idx, timestamp, metadata={"speech": True}))
|
|
|
|
self.assertEqual([item.frame_id for item in vad.read_available()], [1, 2, 3])
|
|
self.assertEqual([item.frame_id for item in stt.read_available()], [1, 2, 3])
|
|
self.assertEqual(vad.read_available(), ())
|
|
self.assertEqual(stt.read_available(), ())
|
|
|
|
def test_audio_hub_processes_capture_before_processed_subscribers_read_it(self) -> None:
|
|
hub = AudioHub(
|
|
processor=FakeWebRtcAudioProcessingProvider(sample_rate=16000, channels=1),
|
|
capture_capacity_ms=100,
|
|
render_capacity_ms=100,
|
|
)
|
|
hub.accept_render(frame(1, 0, metadata={"assistant_audio": True}))
|
|
subscription = hub.subscribe("processed_capture", name="interrupt")
|
|
|
|
processed = hub.accept_capture(frame(2, 20, metadata={"assistant_echo": True, "speech": True}))
|
|
|
|
self.assertTrue(processed.metadata["echo_suppressed"])
|
|
self.assertEqual(subscription.read_available(), (processed,))
|
|
self.assertEqual([item.frame_id for item in hub.render_reference.frames()], [1])
|
|
|
|
def test_audio_hub_reports_ring_and_subscriber_overrun(self) -> None:
|
|
hub = AudioHub(
|
|
processor=FakeWebRtcAudioProcessingProvider(sample_rate=16000, channels=1),
|
|
capture_capacity_ms=40,
|
|
)
|
|
subscription = hub.subscribe("processed_capture", name="slow-stt")
|
|
hub.accept_capture(frame(1, 0))
|
|
self.assertEqual([item.frame_id for item in subscription.read_available()], [1])
|
|
|
|
hub.accept_capture(frame(2, 20))
|
|
hub.accept_capture(frame(3, 40))
|
|
hub.accept_capture(frame(4, 60))
|
|
|
|
self.assertEqual([item.frame_id for item in subscription.read_available()], [3, 4])
|
|
self.assertGreaterEqual(subscription.missed_frames, 1)
|
|
self.assertTrue(any(item.code == ErrorCode.AUDIO_BUFFER_OVERRUN for item in hub.diagnostics))
|
|
self.assertTrue(any(item.subscriber == "slow-stt" for item in hub.diagnostics))
|
|
|
|
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_detected_but_unwired_webrtc_provider_fails_when_required(self) -> None:
|
|
with patch("owner_voice_pet.full_duplex_audio.webrtc_apm_probe") as probe:
|
|
probe.return_value = type(
|
|
"Health",
|
|
(),
|
|
{"available": True, "message": "found 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_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()
|