[WebRTC音频底座]:完成全双工音频基础模块,包含环形缓冲、APM接口和fake回声抑制测试
This commit is contained in:
@@ -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. 全双工状态机、事件总线与取消机制
|
||||
|
||||
|
||||
@@ -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,
|
||||
)
|
||||
@@ -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:
|
||||
|
||||
@@ -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()
|
||||
Reference in New Issue
Block a user