[WebRTC音频底座]:完成全双工音频基础模块,包含环形缓冲、APM接口和fake回声抑制测试

This commit is contained in:
mkbk
2026-06-18 21:44:55 +08:00
parent 9fd8eaf7eb
commit 8b3ffe0ef3
4 changed files with 441 additions and 7 deletions
@@ -8,13 +8,13 @@
## 2. WebRTC APM 与音频环形缓冲 ## 2. WebRTC APM 与音频环形缓冲
- [ ] 2.1 设计 `AudioFrame` 数据结构;前置条件:内部采样率和帧长方案已确认;优先级:P0;验收标准:包含 samples、sample_rate、channels、timestamp、frame_id;测试要点:frame serialization fake fixture 可稳定回放。 - [x] 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` - [x] 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。 - [x] 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。 - [x] 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。 - [x] 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 状态。 - [x] 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.7 增加 APM fallback 策略;前置条件:fake/真实 provider 接口完成;优先级:P1;验收标准:`OWNER_AUDIO_APM_REQUIRED` 控制失败即退或带标记降级;测试要点:APM 不可用时不进入假全双工。
## 3. 全双工状态机、事件总线与取消机制 ## 3. 全双工状态机、事件总线与取消机制
+255
View File
@@ -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,
)
+36
View File
@@ -23,6 +23,10 @@ class ErrorCode(str, Enum):
AUDIO_PERMISSION_DENIED = "AUDIO_PERMISSION_DENIED" AUDIO_PERMISSION_DENIED = "AUDIO_PERMISSION_DENIED"
AUDIO_STREAM_UNDERRUN = "AUDIO_STREAM_UNDERRUN" AUDIO_STREAM_UNDERRUN = "AUDIO_STREAM_UNDERRUN"
AUDIO_FORMAT_UNSUPPORTED = "AUDIO_FORMAT_UNSUPPORTED" 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" CONFIG_MISSING_VALUE = "CONFIG_MISSING_VALUE"
WAKE_MODEL_MISSING = "WAKE_MODEL_MISSING" WAKE_MODEL_MISSING = "WAKE_MODEL_MISSING"
WAKE_MODEL_LOAD_FAILED = "WAKE_MODEL_LOAD_FAILED" WAKE_MODEL_LOAD_FAILED = "WAKE_MODEL_LOAD_FAILED"
@@ -78,6 +82,38 @@ class AudioFrame:
if self.frame_id < 0: if self.frame_id < 0:
raise ValueError("frame_id must be non-negative") 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) @dataclass(frozen=True, slots=True)
class AudioSegment: class AudioSegment:
+143
View File
@@ -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()