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.
202 lines
7.0 KiB
Python
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
|