- auth from Authorization/Device-Id headers (firmware v2.5.0), frame fallback kept
- server hello carries transport=websocket + session_id (firmware ParseServerHello)
- /ota GET+POST returns {websocket:{url,token,version}, server_time:{timestamp,timezone_offset}, firmware:{version}}
- config.yaml: device registered by MAC, ota block added
- 21 tests pass (added OTA + header-auth coverage)
166 lines
5.5 KiB
Python
166 lines
5.5 KiB
Python
"""AC-003: handshake + auth + chunker + barge-in + latency + profile + e2e turn."""
|
|
from __future__ import annotations
|
|
|
|
import asyncio
|
|
|
|
import pytest
|
|
|
|
from app.config import Config
|
|
from app.xiaozhi.protocol import build_hello, build_server_hello
|
|
from app.xiaozhi import messages as msg
|
|
from app.xiaozhi.session import Session
|
|
from app.stt import make_stt
|
|
from app.tts import make_tts, chunk_text, TTSQueue
|
|
from app.hermes import HermesClient, VoiceProfile
|
|
from app.audio.vad import make_vad, EnergyVAD
|
|
from app.security import DeviceAuth, AuthError
|
|
from app.metrics import LatencyLogger
|
|
from app.gpu import GpuResourceManager, GpuMonitor, GpuState, ResourceMode
|
|
|
|
|
|
# ---- handshake (REQ-001) -------------------------------------------------
|
|
def test_hello_parses_firmware_aliases():
|
|
hello = build_hello({"deviceId": "d1", "authorization": "tok", "audio": {"sample_rate": 16000}})
|
|
assert hello.device_id == "d1"
|
|
assert hello.authorization == "tok"
|
|
assert hello.audio_params.sample_rate == 16000
|
|
|
|
|
|
def test_server_hello_firmware_compat():
|
|
hello = build_hello({"device_id": "d1"})
|
|
reply = build_server_hello(hello)
|
|
# FIRMWARE: ParseServerHello requires transport + session_id
|
|
assert reply.type == "hello"
|
|
assert reply.transport == "websocket"
|
|
assert reply.session_id # non-empty uuid
|
|
assert reply.audio_params.sample_rate == 16000
|
|
|
|
|
|
# ---- auth (REQ-009) ------------------------------------------------------
|
|
def test_auth_ok_and_bad_token():
|
|
auth = DeviceAuth({"d1": DeviceAuth.hash_token("secret")})
|
|
assert auth.authenticate("d1", "secret") is True
|
|
with pytest.raises(AuthError):
|
|
auth.authenticate("d1", "wrong")
|
|
|
|
|
|
def test_auth_unknown_device():
|
|
auth = DeviceAuth({"d1": "x"})
|
|
with pytest.raises(AuthError):
|
|
auth.authenticate("nobody", "secret")
|
|
|
|
|
|
def test_auth_disabled_allows():
|
|
auth = DeviceAuth({}, enabled=False)
|
|
assert auth.authenticate("anything", "whatever") is True
|
|
|
|
|
|
# ---- chunker (REQ-006) ---------------------------------------------------
|
|
def test_chunker_splits_on_sentence():
|
|
assert chunk_text("สวัสดี. แล้วไง? ครับ") == ["สวัสดี.", "แล้วไง?", "ครับ"]
|
|
|
|
|
|
def test_chunker_empty():
|
|
assert chunk_text("") == []
|
|
assert chunk_text(" ") == []
|
|
|
|
|
|
def test_chunker_caps_long_fragment():
|
|
long = "คำ" * 100 # > 60 chars, no punctuation
|
|
chunks = chunk_text(long)
|
|
assert all(len(c) <= 60 for c in chunks)
|
|
|
|
|
|
# ---- barge-in queue (REQ-007) --------------------------------------------
|
|
@pytest.mark.asyncio
|
|
async def test_tts_queue_cancel_drops_pending():
|
|
q = TTSQueue()
|
|
for i in range(5):
|
|
await q.put(f"c{i}".encode())
|
|
dropped = q.cancel()
|
|
assert dropped == 5
|
|
assert q.cancelled is True
|
|
|
|
|
|
# ---- latency metric (REQ-011) --------------------------------------------
|
|
def test_latency_eos_to_first_audio():
|
|
import time
|
|
|
|
logger = LatencyLogger()
|
|
base = time.monotonic()
|
|
logger.speech_end()
|
|
time.sleep(0)
|
|
logger.first_audio_packet()
|
|
delta = logger.eos_to_first_audio()
|
|
assert delta is not None and delta >= 0
|
|
|
|
|
|
# ---- voice profile (REQ-004 / AC-004) ------------------------------------
|
|
def test_voice_profile_forces_none_and_search():
|
|
p = VoiceProfile("none", ["session_search"])
|
|
assert p.reasoning == "none"
|
|
assert "session_search" in p.tools
|
|
|
|
|
|
def test_voice_profile_rejects_exec_tools():
|
|
with pytest.raises(AssertionError):
|
|
VoiceProfile("none", ["session_search", "terminal"])
|
|
|
|
|
|
def test_voice_profile_rejects_reasoning():
|
|
with pytest.raises(AssertionError):
|
|
VoiceProfile("medium", ["session_search"])
|
|
|
|
|
|
# ---- GPU modes (REQ-010) -------------------------------------------------
|
|
def test_gpu_plans():
|
|
rm = GpuResourceManager(GpuMonitor(GpuState(vram_total_gb=16, vram_free_gb=16)))
|
|
assert rm.plan(ResourceMode.FULL) == {"stt": True, "tts": True, "unload_tts": False}
|
|
assert rm.plan(ResourceMode.LITE)["tts"] is False
|
|
assert rm.plan(ResourceMode.OFF)["stt"] is False
|
|
assert rm.can_fit(ResourceMode.FULL) is True
|
|
# 1GB free can't fit FULL (6GB needed)
|
|
rm2 = GpuResourceManager(GpuMonitor(GpuState(vram_total_gb=16, vram_free_gb=1)))
|
|
assert rm2.can_fit(ResourceMode.FULL) is False
|
|
|
|
|
|
# ---- VAD end-of-speech (REQ-002) -----------------------------------------
|
|
def test_vad_detects_speech_vs_silence():
|
|
vad = EnergyVAD(threshold=300.0)
|
|
loud = b"\xff\x7f" * 200 # 0x7FFF = 32767, rms >> threshold
|
|
silent = b"\x00\x00" * 200
|
|
assert vad.is_speech(loud) is True
|
|
assert vad.is_speech(silent) is False
|
|
|
|
|
|
# ---- end-to-end spoken turn (AC-005) -------------------------------------
|
|
@pytest.mark.asyncio
|
|
async def test_full_spoken_turn_mocks():
|
|
cfg = Config.from_dict(
|
|
{
|
|
"stt": {"engine": "mock"},
|
|
"tts": {"engine": "mock"},
|
|
"hermes": {"transport": "mock", "reasoning": "none", "tools": ["session_search"]},
|
|
}
|
|
)
|
|
profile = VoiceProfile(cfg.hermes.reasoning, cfg.hermes.tools)
|
|
client = HermesClient(cfg.hermes, profile)
|
|
session = Session(
|
|
client=client,
|
|
stt=make_stt(cfg.stt),
|
|
tts=make_tts(cfg.tts),
|
|
vad=make_vad("energy"),
|
|
logger=LatencyLogger(),
|
|
)
|
|
|
|
events = []
|
|
async for kind, payload in session.run_turn(b"\x00\x00" * 100):
|
|
events.append((kind, payload))
|
|
|
|
kinds = [k for k, _ in events]
|
|
assert "stt" in kinds
|
|
assert "tts_text" in kinds
|
|
assert "audio" in kinds
|
|
# History recorded the turn (multi-turn / REQ-005).
|
|
assert len(client.session().history) == 1
|