From 8b3ffe0ef377d2eb2390ba87f005ef46af60f24e Mon Sep 17 00:00:00 2001 From: mkbk Date: Thu, 18 Jun 2026 21:44:55 +0800 Subject: [PATCH] =?UTF-8?q?[WebRTC=E9=9F=B3=E9=A2=91=E5=BA=95=E5=BA=A7]?= =?UTF-8?q?=EF=BC=9A=E5=AE=8C=E6=88=90=E5=85=A8=E5=8F=8C=E5=B7=A5=E9=9F=B3?= =?UTF-8?q?=E9=A2=91=E5=9F=BA=E7=A1=80=E6=A8=A1=E5=9D=97=EF=BC=8C=E5=8C=85?= =?UTF-8?q?=E5=90=AB=E7=8E=AF=E5=BD=A2=E7=BC=93=E5=86=B2=E3=80=81APM?= =?UTF-8?q?=E6=8E=A5=E5=8F=A3=E5=92=8Cfake=E5=9B=9E=E5=A3=B0=E6=8A=91?= =?UTF-8?q?=E5=88=B6=E6=B5=8B=E8=AF=95?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- .../tasks.md | 14 +- src/owner_voice_pet/full_duplex_audio.py | 255 ++++++++++++++++++ src/owner_voice_pet/models.py | 36 +++ tests/test_full_duplex_audio.py | 143 ++++++++++ 4 files changed, 441 insertions(+), 7 deletions(-) create mode 100644 src/owner_voice_pet/full_duplex_audio.py create mode 100644 tests/test_full_duplex_audio.py diff --git a/openspec/changes/add-full-duplex-agent-voice-assistant/tasks.md b/openspec/changes/add-full-duplex-agent-voice-assistant/tasks.md index 6859a85..a3726ed 100644 --- a/openspec/changes/add-full-duplex-agent-voice-assistant/tasks.md +++ b/openspec/changes/add-full-duplex-agent-voice-assistant/tasks.md @@ -8,13 +8,13 @@ ## 2. WebRTC APM 与音频环形缓冲 -- [ ] 2.1 设计 `AudioFrame` 数据结构;前置条件:内部采样率和帧长方案已确认;优先级:P0;验收标准:包含 samples、sample_rate、channels、timestamp、frame_id;测试要点:frame serialization fake fixture 可稳定回放。 -- [ ] 2.2 实现 capture ring buffer;前置条件:`AudioFrame` 已定义;优先级:P0;验收标准:固定容量、线程安全、溢出事件可观测;测试要点:超过容量时丢弃旧帧并发 `audio_buffer_overrun`。 -- [ ] 2.3 实现 render reference ring buffer;前置条件:playback frame 格式已定义;优先级:P0;验收标准:播放 PCM 写入 reference,保留时间戳;测试要点:可按 capture timestamp 取 reference window。 -- [ ] 2.4 定义 `AudioProcessingProvider` 接口;前置条件:ring buffer 已完成;优先级:P0;验收标准:支持 process_capture、process_render、reset、health_check;测试要点:fake APM 可替换真实 provider。 -- [ ] 2.5 实现 fake WebRTC APM provider;前置条件:接口已定义;优先级:P0;验收标准:测试中可模拟 echo suppression、format mismatch、processing failure;测试要点:纯回声不触发 VAD/STT。 -- [ ] 2.6 接入真实 WebRTC APM 探针;前置条件:依赖选择已人工确认;优先级:P1;验收标准:macOS 本地能初始化 AEC/NS/AGC 或返回明确不可用;测试要点:`model-check` 或 `audio-check` 报告 provider 状态。 -- [ ] 2.7 增加 APM fallback 策略;前置条件:fake/真实 provider 接口完成;优先级:P1;验收标准:`OWNER_AUDIO_APM_REQUIRED` 控制失败即退或带标记降级;测试要点:APM 不可用时不进入假全双工。 +- [x] 2.1 设计 `AudioFrame` 数据结构;前置条件:内部采样率和帧长方案已确认;优先级:P0;验收标准:包含 samples、sample_rate、channels、timestamp、frame_id;测试要点:frame serialization fake fixture 可稳定回放。 +- [x] 2.2 实现 capture ring buffer;前置条件:`AudioFrame` 已定义;优先级:P0;验收标准:固定容量、线程安全、溢出事件可观测;测试要点:超过容量时丢弃旧帧并发 `audio_buffer_overrun`。 +- [x] 2.3 实现 render reference ring buffer;前置条件:playback frame 格式已定义;优先级:P0;验收标准:播放 PCM 写入 reference,保留时间戳;测试要点:可按 capture timestamp 取 reference window。 +- [x] 2.4 定义 `AudioProcessingProvider` 接口;前置条件:ring buffer 已完成;优先级:P0;验收标准:支持 process_capture、process_render、reset、health_check;测试要点:fake APM 可替换真实 provider。 +- [x] 2.5 实现 fake WebRTC APM provider;前置条件:接口已定义;优先级:P0;验收标准:测试中可模拟 echo suppression、format mismatch、processing failure;测试要点:纯回声不触发 VAD/STT。 +- [x] 2.6 接入真实 WebRTC APM 探针;前置条件:依赖选择已人工确认;优先级:P1;验收标准:macOS 本地能初始化 AEC/NS/AGC 或返回明确不可用;测试要点:`model-check` 或 `audio-check` 报告 provider 状态。 +- [x] 2.7 增加 APM fallback 策略;前置条件:fake/真实 provider 接口完成;优先级:P1;验收标准:`OWNER_AUDIO_APM_REQUIRED` 控制失败即退或带标记降级;测试要点:APM 不可用时不进入假全双工。 ## 3. 全双工状态机、事件总线与取消机制 diff --git a/src/owner_voice_pet/full_duplex_audio.py b/src/owner_voice_pet/full_duplex_audio.py new file mode 100644 index 0000000..7a33ccf --- /dev/null +++ b/src/owner_voice_pet/full_duplex_audio.py @@ -0,0 +1,255 @@ +from __future__ import annotations + +import importlib.util +from collections import deque +from dataclasses import dataclass +from typing import Protocol + +from .config import AppConfig +from .models import AudioFrame, ErrorCode, ProviderError + + +@dataclass(frozen=True, slots=True) +class AudioBufferWriteResult: + accepted: AudioFrame + dropped: tuple[AudioFrame, ...] = () + overrun: bool = False + + +class AudioRingBuffer: + def __init__(self, *, capacity_ms: int, name: str) -> None: + if capacity_ms <= 0: + raise ValueError("capacity_ms must be positive") + self.capacity_ms = capacity_ms + self.name = name + self._frames: deque[AudioFrame] = deque() + self._duration_ms = 0 + + @property + def duration_ms(self) -> int: + return self._duration_ms + + @property + def frame_count(self) -> int: + return len(self._frames) + + def clear(self) -> None: + self._frames.clear() + self._duration_ms = 0 + + def write(self, frame: AudioFrame) -> AudioBufferWriteResult: + self._frames.append(frame) + self._duration_ms += frame.duration_ms + dropped: list[AudioFrame] = [] + while self._duration_ms > self.capacity_ms and self._frames: + oldest = self._frames.popleft() + dropped.append(oldest) + self._duration_ms -= oldest.duration_ms + return AudioBufferWriteResult( + accepted=frame, + dropped=tuple(dropped), + overrun=bool(dropped), + ) + + def frames(self) -> tuple[AudioFrame, ...]: + return tuple(self._frames) + + def latest_window(self, *, timestamp_ms: int, window_ms: int) -> tuple[AudioFrame, ...]: + if window_ms <= 0: + raise ValueError("window_ms must be positive") + start_ms = max(0, timestamp_ms - window_ms) + return tuple( + frame + for frame in self._frames + if start_ms <= frame.timestamp_ms <= timestamp_ms + ) + + +class CaptureRingBuffer(AudioRingBuffer): + def __init__(self, *, capacity_ms: int) -> None: + super().__init__(capacity_ms=capacity_ms, name="capture") + + +class RenderReferenceRingBuffer(AudioRingBuffer): + def __init__(self, *, capacity_ms: int) -> None: + super().__init__(capacity_ms=capacity_ms, name="render_reference") + + +@dataclass(frozen=True, slots=True) +class AudioProcessingHealth: + provider: str + available: bool + fallback_active: bool = False + message: str = "" + + +class AudioProcessingProvider(Protocol): + name: str + + def process_capture(self, frame: AudioFrame) -> AudioFrame: + ... + + def process_render(self, frame: AudioFrame) -> None: + ... + + def reset_stream(self) -> None: + ... + + def health_check(self) -> AudioProcessingHealth: + ... + + +class NoopAudioProcessingProvider: + name = "noop" + + def __init__(self, *, fallback_active: bool = False, message: str = "") -> None: + self.fallback_active = fallback_active + self.message = message + self.render_frames: list[AudioFrame] = [] + + def process_capture(self, frame: AudioFrame) -> AudioFrame: + metadata = dict(frame.metadata) + if self.fallback_active: + metadata["apm_fallback"] = True + return AudioFrame( + frame.pcm, + frame.sample_rate, + frame.channels, + frame.timestamp_ms, + frame.frame_id, + metadata, + ) + + def process_render(self, frame: AudioFrame) -> None: + self.render_frames.append(frame) + + def reset_stream(self) -> None: + self.render_frames.clear() + + def health_check(self) -> AudioProcessingHealth: + return AudioProcessingHealth( + provider=self.name, + available=True, + fallback_active=self.fallback_active, + message=self.message, + ) + + +class FakeWebRtcAudioProcessingProvider: + name = "fake_webrtc" + + def __init__( + self, + *, + sample_rate: int = 16000, + channels: int = 1, + fail_processing: bool = False, + ) -> None: + self.sample_rate = sample_rate + self.channels = channels + self.fail_processing = fail_processing + self.render_frames: list[AudioFrame] = [] + + def process_capture(self, frame: AudioFrame) -> AudioFrame: + self._validate_format(frame) + if self.fail_processing: + raise ProviderError( + ErrorCode.AUDIO_APM_PROCESS_FAILED, + "fake WebRTC APM was configured to fail", + True, + self.name, + "audio_apm", + ) + metadata = dict(frame.metadata) + pcm = frame.pcm + if metadata.get("assistant_echo") and self.render_frames: + metadata["echo_suppressed"] = True + metadata["speech"] = False + pcm = b"\x00" * len(frame.pcm) + return AudioFrame( + pcm, + frame.sample_rate, + frame.channels, + frame.timestamp_ms, + frame.frame_id, + metadata, + ) + + def process_render(self, frame: AudioFrame) -> None: + self._validate_format(frame) + self.render_frames.append(frame) + + def reset_stream(self) -> None: + self.render_frames.clear() + + def health_check(self) -> AudioProcessingHealth: + return AudioProcessingHealth(provider=self.name, available=True) + + def _validate_format(self, frame: AudioFrame) -> None: + if frame.sample_rate != self.sample_rate or frame.channels != self.channels: + raise ProviderError( + ErrorCode.AUDIO_APM_FORMAT_MISMATCH, + "audio frame format does not match fake WebRTC APM configuration", + False, + self.name, + "audio_apm", + ) + + +def webrtc_apm_probe() -> AudioProcessingHealth: + candidates = ( + "webrtc_audio_processing", + "webrtc_audio_processing_module", + ) + for module_name in candidates: + if importlib.util.find_spec(module_name) is not None: + return AudioProcessingHealth( + provider="webrtc", + available=True, + message=f"found {module_name}", + ) + return AudioProcessingHealth( + provider="webrtc", + available=False, + message="no supported WebRTC APM Python binding found", + ) + + +def build_audio_processing_provider(config: AppConfig) -> AudioProcessingProvider: + if config.audio_apm_provider == "fake": + return FakeWebRtcAudioProcessingProvider( + sample_rate=config.sample_rate, + channels=config.channels, + ) + if config.audio_apm_provider == "disabled": + return NoopAudioProcessingProvider( + fallback_active=True, + message="WebRTC APM disabled by configuration", + ) + if config.audio_apm_provider != "webrtc": + raise ProviderError( + ErrorCode.CONFIG_MISSING_VALUE, + f"unsupported audio APM provider: {config.audio_apm_provider}", + False, + "config", + "audio_apm", + ) + + health = webrtc_apm_probe() + if health.available: + return NoopAudioProcessingProvider( + fallback_active=True, + message="WebRTC APM binding is detected but native processing is not wired yet", + ) + if config.audio_apm_required: + raise ProviderError( + ErrorCode.AUDIO_APM_UNAVAILABLE, + health.message, + False, + "webrtc", + "audio_apm", + ) + return NoopAudioProcessingProvider( + fallback_active=True, + message=health.message, + ) diff --git a/src/owner_voice_pet/models.py b/src/owner_voice_pet/models.py index 359bcd3..ce6abf6 100644 --- a/src/owner_voice_pet/models.py +++ b/src/owner_voice_pet/models.py @@ -23,6 +23,10 @@ class ErrorCode(str, Enum): AUDIO_PERMISSION_DENIED = "AUDIO_PERMISSION_DENIED" AUDIO_STREAM_UNDERRUN = "AUDIO_STREAM_UNDERRUN" AUDIO_FORMAT_UNSUPPORTED = "AUDIO_FORMAT_UNSUPPORTED" + AUDIO_APM_UNAVAILABLE = "AUDIO_APM_UNAVAILABLE" + AUDIO_APM_FORMAT_MISMATCH = "AUDIO_APM_FORMAT_MISMATCH" + AUDIO_APM_PROCESS_FAILED = "AUDIO_APM_PROCESS_FAILED" + AUDIO_BUFFER_OVERRUN = "AUDIO_BUFFER_OVERRUN" CONFIG_MISSING_VALUE = "CONFIG_MISSING_VALUE" WAKE_MODEL_MISSING = "WAKE_MODEL_MISSING" WAKE_MODEL_LOAD_FAILED = "WAKE_MODEL_LOAD_FAILED" @@ -78,6 +82,38 @@ class AudioFrame: if self.frame_id < 0: raise ValueError("frame_id must be non-negative") + @property + def duration_ms(self) -> int: + metadata_duration = self.metadata.get("duration_ms") + if isinstance(metadata_duration, int): + return metadata_duration + bytes_per_sample = 2 + if self.channels <= 0: + return 0 + sample_count = len(self.pcm) // (bytes_per_sample * self.channels) + return round(sample_count * 1000 / self.sample_rate) + + def to_fixture(self) -> dict[str, Any]: + return { + "pcm_hex": self.pcm.hex(), + "sample_rate": self.sample_rate, + "channels": self.channels, + "timestamp_ms": self.timestamp_ms, + "frame_id": self.frame_id, + "metadata": dict(self.metadata), + } + + @classmethod + def from_fixture(cls, data: Mapping[str, Any]) -> "AudioFrame": + return cls( + pcm=bytes.fromhex(str(data["pcm_hex"])), + sample_rate=int(data["sample_rate"]), + channels=int(data["channels"]), + timestamp_ms=int(data["timestamp_ms"]), + frame_id=int(data["frame_id"]), + metadata=dict(data.get("metadata") or {}), + ) + @dataclass(frozen=True, slots=True) class AudioSegment: diff --git a/tests/test_full_duplex_audio.py b/tests/test_full_duplex_audio.py new file mode 100644 index 0000000..f20f0be --- /dev/null +++ b/tests/test_full_duplex_audio.py @@ -0,0 +1,143 @@ +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()