Files
esp32-server/app/xiaozhi/websocket.py
kunthawat 261b0f3e91 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.
2026-10-03 20:35:38 +07:00

202 lines
7.0 KiB
Python

"""Device-facing WebSocket loop (REQ-001/002/007/009).
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
from ..config import Config
from ..hermes import HermesClient, VoiceProfile
from ..stt import make_stt
from ..tts import make_tts
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
from . import messages as msg
from .session import Session
def _bearer(value: str | None) -> str | None:
"""Strip an optional ``Bearer `` prefix from an Authorization header."""
if value is None:
return None
v = value.strip()
if not v:
return None
if v.lower().startswith("bearer "):
return v[7:].strip()
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,
cfg: Config,
profile: VoiceProfile,
auth: DeviceAuth,
headers: dict[str, str] | None = None,
) -> None:
self.cfg = cfg
self.profile = profile
self.auth = auth
self.headers = headers or {}
async def run(self, ws: WebSocket) -> None:
await ws.accept()
hello = await self._handshake(ws)
if hello is None:
return
client = HermesClient(self.cfg.hermes, self.profile)
session = Session(
client=client,
stt=make_stt(self.cfg.stt),
tts=make_tts(self.cfg.tts),
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 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)
hello = build_hello(json.loads(raw))
except asyncio.TimeoutError:
await ws.close()
return None
except Exception as exc: # noqa: BLE001
await ws.send_text(json.dumps({"type": "error", "message": str(exc)}))
return None
try:
# FIRMWARE: auth from HTTP headers (Bearer token), not the frame.
# Fall back to frame fields for synthetic/test clients.
device_id = self.headers.get("device_id") or hello.device_id
token = _bearer(self.headers.get("authorization"))
if token is None:
token = hello.authorization or ""
self.auth.authenticate(device_id, token)
except AuthError as exc:
await ws.send_text(json.dumps({"type": "error", "message": str(exc)}))
await ws.close()
return None
session_id = str(uuid.uuid4())
await ws.send_text(json.dumps(build_server_hello(hello).model_dump()))
return hello
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:
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.tts_sentence_start(payload)))
elif kind == "audio":
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