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:
@@ -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}")
|
||||
|
||||
@@ -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":
|
||||
|
||||
@@ -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":
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user