feat: wire Opus decode + VAD turn detection to firmware-compatible reply

Device streams continuous 20ms Opus frames with no turn markers. The loop
now decodes Opus, feeds the VAD, fires one turn on end-of-speech, and
replies with the firmware JSON contract (tts.state / stt.text) + Opus audio.

- app/audio/opus.py: RealOpus/Opus16kEncoder/PassthroughOpus; fix add_dll_directory
  to use absolute paths (WinError 87)
- app/xiaozhi/websocket.py: DeviceLoop.run() = decode -> VAD -> run_turn -> reply
- app/config.py: AudioConfig (opus kind + VAD thresholds)
- config.yaml: audio.opus=real, VAD 5/8/3000 frames
- tests/test_app.py: utterance = loud burst + silence; asserts stt + audio + tts stop
- requirements.txt: opuslib>=3

21/21 tests; live handshake verified via zhi.moreminimore.com on venv python.
This commit is contained in:
2026-10-03 20:35:38 +07:00
parent 3bfeff29bd
commit 261b0f3e91
7 changed files with 293 additions and 40 deletions

View File

@@ -1,48 +1,149 @@
"""Opus encode/decode — lazy (CON-002).
Only the codec the device uses is needed; we default to a passthrough so the
pipeline runs on PCM in tests. The real opus (``opuslib`` / ``PyNaCl``-free
``opus``) is imported at construction time, never at module import.
Device-facing codec for the Xiaozhi bridge:
* :class:`OpusCodec` — stateless packet codec. ``decode`` turns one device
Opus packet into PCM; ``encode`` is NOT used for the reply path (use
:class:`Opus16kEncoder`, which owns frame alignment).
* :class:`Opus16kEncoder` — streaming encoder for reply audio: feed arbitrary
16 kHz mono PCM, get back exact 20 ms (320-sample) Opus frames; ``flush``
emits one final padded frame (smallest valid Opus frame size) so short
utterances still produce a playable packet.
* :class:`PassthroughOpus` — PCM in, PCM out. Lets the whole pipeline run
without a native codec (tests / synthetic clients).
``RealOpus`` is constructed only when selected (kind ``"real"``); the native
``opus`` library is located via ``OPUS_DLL_DIR`` (env), then a per-venv
``.venv/Scripts/opus`` directory (this repo keeps ``opus.dll`` there).
"""
from __future__ import annotations
import os
from abc import ABC, abstractmethod
def _ensure_native() -> None:
"""Make the native ``opus`` library discoverable by opuslib.
opuslib loads it through ``ctypes.util.find_library('opus')`` (PATH on
Windows). We do NOT assume any particular install layout: first an
explicit ``OPUS_DLL_DIR`` env override, then the venv-local dir next to
this repo (``.venv/Scripts/opus``), which this project ships the DLL in.
"""
import ctypes.util
if ctypes.util.find_library("opus"):
return
repo_root = os.path.dirname(
os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
)
candidates = [os.environ.get("OPUS_DLL_DIR", "")]
candidates.append(os.path.join(os.getcwd(), ".venv", "Scripts", "opus"))
candidates.append(os.path.join(repo_root, ".venv", "Scripts", "opus"))
# Windows os.add_dll_directory requires an ABSOLUTE path (a relative one
# raises WinError 87); normalize every candidate before use.
for cand in (os.path.abspath(c) for c in candidates if c):
if os.path.isfile(os.path.join(cand, "opus.dll")):
os.environ["PATH"] = cand + os.pathsep + os.environ.get("PATH", "")
if hasattr(os, "add_dll_directory"):
os.add_dll_directory(cand)
break
class OpusCodec(ABC):
@abstractmethod
def encode(self, pcm: bytes) -> bytes: ...
@abstractmethod
def decode(self, data: bytes) -> bytes: ...
def decode(self, data: bytes) -> bytes:
"""One Opus packet -> PCM (16-bit LE mono)."""
class PassthroughOpus(OpusCodec):
"""PCM in, PCM out. Lets the whole pipeline run without a native codec."""
def encode(self, pcm: bytes) -> bytes:
return pcm
def decode(self, data: bytes) -> bytes:
return data
class RealOpus(OpusCodec):
"""Lazy real codec. Fails to construct only if opus is not installed —
acceptable on the Mac; the voice-server has it."""
"""Real Opus codec via ``opuslib``.
def __init__(self, sample_rate: int = 16000, channels: int = 1, frame_ms: int = 20):
import opus # type: ignore # noqa: F401
Frame math (verified): 20 ms @ 16 kHz mono = 320 samples = 640 B PCM.
``Decoder.decode(packet, 320)``; an undecodable packet (loss) returns
silence of the requested frame size — we raise so the caller can skip it
rather than inject silent PCM into the VAD buffer.
"""
self._enc = opus.Encoder(sample_rate, channels, opus.APPLICATION_VOIP)
self._dec = opus.Decoder(sample_rate, channels)
self._frame_ms = frame_ms
FRAME_SAMPLES = 320 # 20 ms @ 16 kHz
def encode(self, pcm: bytes) -> bytes:
return self._enc.encode(pcm, self._frame_ms)
def __init__(self, sample_rate: int = 16000, channels: int = 1) -> None:
_ensure_native()
from opuslib.classes import Decoder # lazy (CON-002)
self._dec = Decoder(sample_rate, channels)
def decode(self, data: bytes) -> bytes:
return self._dec.decode(data, 20000)
pcm = self._dec.decode(data, self.FRAME_SAMPLES)
if pcm is None:
raise ValueError("undecodable opus packet")
return pcm
class Opus16kEncoder:
"""Streaming 16 kHz mono PCM -> exact 20 ms Opus frames.
Feed PCM in any size; the encoder buffers a partial frame and emits
complete 320-sample frames as they become available. ``flush`` pads the
remainder up to the smallest valid Opus frame (2.5 ms / 40 samples) so
the device always gets a decodable final packet.
"""
FRAME_SAMPLES = 320 # 20 ms @ 16 kHz
FRAME_BYTES = 640
# Valid Opus frame sizes (2.5/5/10/20/40/60 ms) in samples @ 16 kHz.
_VALID = (40, 80, 160, 320, 640, 960)
def __init__(self, sample_rate: int = 16000, channels: int = 1) -> None:
_ensure_native()
from opuslib.classes import Encoder
self._enc = Encoder(sample_rate, channels, "voip")
self._carry = b""
def feed(self, pcm: bytes) -> list[bytes]:
"""Add PCM; return any complete 20 ms frames as Opus packets."""
self._carry += pcm
out: list[bytes] = []
while len(self._carry) >= self.FRAME_BYTES:
frame = self._carry[: self.FRAME_BYTES]
self._carry = self._carry[self.FRAME_BYTES :]
out.append(self._enc.encode(frame, self.FRAME_SAMPLES))
return out
def flush(self) -> list[bytes]:
"""Emit the trailing partial frame (padded to a valid frame size)."""
n = len(self._carry) // 2
if n == 0:
return []
for valid in self._VALID:
if n <= valid:
target = valid
break
else: # pragma: no cover — carry is always < 320 samples
target = 960
pad = (target - n) * 2
pcm = self._carry + b"\x00" * pad
self._carry = b""
return [self._enc.encode(pcm, target)]
class _PassthroughEncoder:
"""Encoder shim with the same interface as :class:`Opus16kEncoder`."""
def feed(self, pcm: bytes) -> list[bytes]:
return [pcm] if pcm else []
def flush(self) -> list[bytes]:
return []
def make_opus(kind: str = "passthrough", **kw) -> OpusCodec:
@@ -51,3 +152,11 @@ def make_opus(kind: str = "passthrough", **kw) -> OpusCodec:
if kind == "real":
return RealOpus(**kw)
raise ValueError(f"unknown opus kind: {kind}")
def make_encoder(kind: str = "passthrough", **kw):
if kind == "passthrough":
return _PassthroughEncoder()
if kind == "real":
return Opus16kEncoder(**kw)
raise ValueError(f"unknown opus kind: {kind}")

View File

@@ -56,6 +56,18 @@ class ProtocolConfig(BaseModel):
frame_ms: int = 20 # typical Xiaozhi Opus frame size
class AudioConfig(BaseModel):
"""Device audio path: the device streams continuous 20 ms Opus frames
(no turn markers), so the bridge decodes Opus, accumulates through a
VAD, and fires a turn on end-of-speech. ``passthrough`` treats frames as
raw PCM (synthetic clients / tests without a native codec)."""
opus: Literal["passthrough", "real"] = "passthrough"
vad_min_speech_frames: int = 5 # 100 ms of speech before an utterance counts
vad_min_silence_frames: int = 8 # 160 ms of silence ends the utterance
vad_max_utterance_frames: int = 3000 # 60 s safety flush
class OtaConfig(BaseModel):
"""OTA endpoint data (GET/POST /ota) — xiaozhi device onboarding.
@@ -95,6 +107,7 @@ class Config(BaseModel):
security: SecurityConfig = Field(default_factory=SecurityConfig)
ota: Optional[OtaConfig] = None
gpu: GpuConfig = Field(default_factory=GpuConfig)
audio: AudioConfig = Field(default_factory=AudioConfig)
@classmethod
def from_dict(cls, data: dict) -> "Config":

View File

@@ -46,6 +46,28 @@ def thinking() -> dict:
return state("thinking")
# ---------------------------------------------------------------------------
# Firmware reply shapes (xiaozhi-esp32 v2.5.0, verified in application.cc):
# the device understands type "stt" (text), "tts" (state start /
# sentence_start + text / stop) and "llm" (emotion). "state"/"text" above
# are kept for synthetic clients that display transcripts.
# ---------------------------------------------------------------------------
def tts_start() -> dict:
return {"type": MsgType.TTS, "state": "start"}
def tts_sentence_start(text: str) -> dict:
return {"type": MsgType.TTS, "state": "sentence_start", "text": text}
def tts_stop() -> dict:
return {"type": MsgType.TTS, "state": "stop"}
def llm(emotion: str) -> dict:
return {"type": "llm", "emotion": emotion}
def build_message(kind: str, text: str = "", **extra) -> dict:
"""Small dispatcher used by the session to emit device-facing messages."""
if kind == "text":

View File

@@ -1,13 +1,19 @@
"""Device-facing WebSocket loop (REQ-001/002/007/009).
Encapsulates the handshake → auth → conversational-turn loop so ``main.py``
just wires a FastAPI endpoint to :func:`run`. One instance per connection.
Firmware audio contract (xiaozhi-esp32 v2.5.0, verified from source):
the device streams continuous 20 ms Opus frames as binary WS frames with
NO per-turn markers. The bridge therefore decodes Opus, accumulates PCM
through a VAD, and fires one spoken turn on end-of-speech. Replies are
binary Opus frames plus the JSON metadata the firmware understands:
``stt.text`` (user transcript), ``tts.state`` = start / sentence_start
(+``text``) / stop, ``llm.emotion``.
"""
from __future__ import annotations
import asyncio
import json
import uuid
from dataclasses import dataclass
from fastapi import WebSocket
@@ -15,7 +21,8 @@ from ..config import Config
from ..hermes import HermesClient, VoiceProfile
from ..stt import make_stt
from ..tts import make_tts
from ..audio.vad import make_vad
from ..audio.vad import VAD, make_vad
from ..audio.opus import make_opus, make_encoder
from ..security import DeviceAuth, AuthError
from ..metrics import LatencyLogger
from .protocol import build_hello, build_server_hello
@@ -35,6 +42,49 @@ def _bearer(value: str | None) -> str | None:
return v
@dataclass
class _TurnDetector:
"""VAD end-of-speech over a continuous frame stream.
``feed`` one frame per call; returns True exactly once when the
utterance ends (silence gap after speech, or the hard length cap).
Frames before speech is detected are ignored (noise floor).
"""
vad: VAD
min_speech_frames: int
min_silence_frames: int
max_utterance_bytes: int
_in_speech: bool = False
_speech_frames: int = 0
_silence_frames: int = 0
_bytes: int = 0
def feed(self, pcm: bytes) -> bool:
self._bytes += len(pcm)
if self.vad.is_speech(pcm):
self._speech_frames += 1
self._silence_frames = 0
else:
self._silence_frames += 1
if not self._in_speech:
self._in_speech = self._speech_frames >= self.min_speech_frames
return False
ended = (
self._silence_frames >= self.min_silence_frames
or self._bytes >= self.max_utterance_bytes
)
if ended:
self.reset()
return ended
def reset(self) -> None:
self._in_speech = False
self._speech_frames = 0
self._silence_frames = 0
self._bytes = 0
class DeviceLoop:
def __init__(
self,
@@ -62,17 +112,42 @@ class DeviceLoop:
vad=make_vad("energy"),
logger=LatencyLogger(),
)
decoder = make_opus(
self.cfg.audio.opus, sample_rate=self.cfg.protocol.opus_sample_rate
)
frame_bytes = self.cfg.protocol.opus_sample_rate // 50 * 2 # 20 ms
detector = _TurnDetector(
session.vad,
min_speech_frames=self.cfg.audio.vad_min_speech_frames,
min_silence_frames=self.cfg.audio.vad_min_silence_frames,
max_utterance_bytes=self.cfg.audio.vad_max_utterance_frames * frame_bytes,
)
try:
while True:
data = await ws.receive()
if data.get("type") == "websocket.disconnect":
break
audio = data.get("bytes")
if audio is not None:
await self._turn(ws, session, audio)
if audio is None:
continue
pcm = self._decode(decoder, audio)
if pcm is None:
continue
session.feed_audio(pcm)
if detector.feed(pcm):
utterance = session.end_of_speech()
if utterance:
await self._turn(ws, session, utterance)
finally:
session.barge_in()
@staticmethod
def _decode(decoder, audio: bytes):
try:
return decoder.decode(audio)
except Exception:
return None
async def _handshake(self, ws: WebSocket):
try:
raw = await asyncio.wait_for(ws.receive_text(), timeout=10)
@@ -99,15 +174,28 @@ class DeviceLoop:
await ws.send_text(json.dumps(build_server_hello(hello).model_dump()))
return hello
async def _turn(self, ws: WebSocket, session: Session, pcm: bytes) -> None:
async def _turn(self, ws: WebSocket, session: Session, user_pcm: bytes) -> None:
"""Run one spoken turn; reply with firmware-compatible frames.
Binary WS frames carry the Opus-encoded reply audio; JSON carries
``stt`` / ``tts`` / ``llm`` metadata. The device keeps streaming
mic audio while we speak (no device-side mute in v2.5.0); frames
arriving during the turn are buffered by the transport and
consumed afterwards (PoC-acceptable).
"""
encoder = make_encoder(self.cfg.audio.opus)
try:
async for kind, payload in session.run_turn(pcm):
await ws.send_text(json.dumps(msg.tts_start()))
async for kind, payload in session.run_turn(user_pcm):
if kind == "stt":
await ws.send_text(json.dumps(msg.stt(payload)))
elif kind == "tts_text":
await ws.send_text(json.dumps(msg.state("speaking")))
await ws.send_text(json.dumps(msg.text(payload)))
await ws.send_text(json.dumps(msg.tts_sentence_start(payload)))
elif kind == "audio":
await ws.send_bytes(payload)
for packet in encoder.feed(payload):
await ws.send_bytes(packet)
for packet in encoder.flush():
await ws.send_bytes(packet)
await ws.send_text(json.dumps(msg.tts_stop()))
except Exception: # keep the device connected on a bad turn
pass

View File

@@ -16,6 +16,12 @@ protocol:
opus_channels: 1 # FIRMWARE
frame_ms: 20 # FIRMWARE
audio:
opus: real # real = decode device Opus | passthrough = raw PCM (tests)
vad_min_speech_frames: 5 # 100 ms of speech before an utterance counts
vad_min_silence_frames: 8 # 160 ms of silence ends the utterance
vad_max_utterance_frames: 3000 # 60 s safety flush
stt:
engine: mock # mock | typhoon | faster-whisper
language: th

View File

@@ -8,6 +8,9 @@ pyyaml>=6.0
httpx>=0.27
websockets>=12
# audio codec (CON-002: lazy-imported; enabled by audio.opus: real)
opuslib>=3
# dev/test
pytest>=8
pytest-asyncio>=0.23

View File

@@ -78,6 +78,8 @@ def test_ws_header_auth_rejects_bad_token():
def test_ws_handshake_and_turn():
"""One spoken utterance (burst of speech frames + trailing silence)
-> firmware-compatible reply: stt text, tts start/sentence/stop, audio."""
client = TestClient(make_app(_cfg()))
with client.websocket_connect("/ws/xiaozhi") as ws:
# hello frame
@@ -85,16 +87,26 @@ def test_ws_handshake_and_turn():
reply = json.loads(ws.receive_text())
assert reply["type"] == "hello"
assert reply["session_id"]
# one utterance -> expect stt + text + audio back
ws.send_bytes(b"\x00\x00" * 64)
got_text = got_audio = False
# one utterance: 6 loud frames (16-bit, RMS >= 300), then silence
# frames that push the VAD past min_silence_frames.
loud = b"\x10\x0f" * 320 # 20 ms @ 16 kHz mono, amplitude 3856
silent = b"\x00\x00" * 320
for _ in range(6):
ws.send_bytes(loud)
for _ in range(12):
try:
kind = ws.receive()
except Exception:
break
if kind.get("text") is not None:
got_text = True
ws.send_bytes(silent)
got_stt = got_audio = False
tts_stop = False
deadline = 0
while not tts_stop and deadline < 256:
frame = ws.receive()
deadline += 1
if frame.get("text") is not None:
payload = json.loads(frame["text"])
if payload.get("type") == "stt":
got_stt = True
if payload.get("type") == "tts" and payload.get("state") == "stop":
tts_stop = True
else:
got_audio = True
assert got_text and got_audio
assert got_stt and got_audio and tts_stop