my_voice_assistant_v2/tests/test_ws.py

157 lines
6.1 KiB
Python
Raw Normal View History

import pytest
from fastapi.testclient import TestClient
from starlette.websockets import WebSocketDisconnect
import app.dependencies as deps
from app.main import app
from app.config import settings
client = TestClient(app)
def _install_stubs(monkeypatch, captured):
class MemLLM:
async def complete(self, text, history=None, session_id=None):
captured["history"] = list(history or [])
return f"Echo: {text}"
class StubTTS:
async def synthesize(self, text, voice=None, audio_format="pcm"):
return b"WSAUDIO"
monkeypatch.setitem(deps.LLM_REGISTRY, "mem", lambda s: MemLLM())
monkeypatch.setitem(deps.TTS_REGISTRY, "stub", lambda s: StubTTS())
return {"llm_provider": "mem", "tts_provider": "stub", "output_endpoint": "loopback"}
def _drain_turn(ws):
assert ws.receive_json()["type"] == "ack"
sem = ws.receive_json()
assert ws.receive_bytes() == b"WSAUDIO"
assert ws.receive_json()["type"] == "done"
return sem
def test_ws_streams_events_and_remembers(monkeypatch):
captured = {}
base = _install_stubs(monkeypatch, captured)
with client.websocket_connect("/ws/chat?session_id=wsconv") as ws:
ws.send_json({"text": "Hallo", **base})
sem1 = _drain_turn(ws)
assert sem1["type"] == "semantic" and sem1["text"] == "Echo: Hallo"
ws.send_json({"text": "Weiter", **base})
_drain_turn(ws)
# Zweiter Turn hat den Verlauf des ersten erhalten.
assert captured["history"] == [
{"role": "user", "content": "Hallo"},
{"role": "assistant", "content": "Echo: Hallo"},
]
def test_ws_unknown_provider_sends_error_event(monkeypatch):
captured = {}
_install_stubs(monkeypatch, captured)
with client.websocket_connect("/ws/chat") as ws:
ws.send_json({"text": "x", "llm_provider": "gibtsnicht"})
err = ws.receive_json()
assert err["type"] == "error" and err["status"] == 422
def test_ws_requires_token_when_auth_enabled(monkeypatch):
monkeypatch.setattr(settings, "auth_enabled", True)
monkeypatch.setattr(settings, "admin_api_key", "k")
with pytest.raises(WebSocketDisconnect):
with client.websocket_connect("/ws/chat"):
pass
def _install_voice_stubs(monkeypatch):
class StubSTT:
async def transcribe(self, audio_bytes, fmt, language=None):
return f"erkannt({len(audio_bytes)})"
class StubLLM:
async def complete(self, text, history=None, session_id=None):
return f"Antwort zu {text}"
class StubTTS:
async def synthesize(self, text, voice=None, audio_format="pcm"):
return b"VOICEAUD"
monkeypatch.setitem(deps.STT_REGISTRY, "ss", lambda s: StubSTT())
monkeypatch.setitem(deps.LLM_REGISTRY, "ll", lambda s: StubLLM())
monkeypatch.setitem(deps.TTS_REGISTRY, "tt", lambda s: StubTTS())
return {
"stt_provider": "ss",
"llm_provider": "ll",
"tts_provider": "tt",
"output_endpoint": "loopback",
}
def test_ws_voice_transcribes_and_answers(monkeypatch):
opts = _install_voice_stubs(monkeypatch)
with client.websocket_connect("/ws/voice?session_id=v1") as ws:
ws.send_bytes(b"PCMDATA") # 7 Bytes
ws.send_bytes(b"MORE") # 4 Bytes -> insgesamt 11
ws.send_json({"type": "end", **opts})
transcript = ws.receive_json()
assert transcript["type"] == "transcript" and transcript["text"] == "erkannt(11)"
assert ws.receive_json()["type"] == "ack"
semantic = ws.receive_json()
assert semantic["type"] == "semantic" and semantic["text"] == "Antwort zu erkannt(11)"
assert ws.receive_bytes() == b"VOICEAUD"
assert ws.receive_json()["type"] == "done"
def test_ws_voice_start_frame_options_are_honored(monkeypatch):
# Provider stehen im START-Frame, der end-Frame ist leer -> muessen trotzdem gelten.
opts = _install_voice_stubs(monkeypatch)
with client.websocket_connect("/ws/voice") as ws:
ws.send_json({"type": "start", "format": "wav", **opts})
ws.send_bytes(b"PCMDATA") # 7 Bytes
ws.send_json({"type": "end"})
transcript = ws.receive_json()
# Stub-STT (aus dem START-Frame) liefert "erkannt(<bytes>)" -> beweist: Override greift
assert transcript["type"] == "transcript" and transcript["text"] == "erkannt(7)"
assert ws.receive_json()["type"] == "ack"
assert ws.receive_json()["type"] == "semantic"
ws.receive_bytes()
assert ws.receive_json()["type"] == "done"
def test_ws_voice_stream_text_emits_token_events(monkeypatch):
# stream:true im START-Frame -> der Server schickt token-Events (Live-Text).
opts = _install_voice_stubs(monkeypatch)
with client.websocket_connect("/ws/voice") as ws:
ws.send_json({"type": "start", "format": "wav", "stream": True, **opts})
ws.send_bytes(b"PCMDATA")
ws.send_json({"type": "end"})
assert ws.receive_json()["type"] == "transcript"
assert ws.receive_json()["type"] == "ack"
assert ws.receive_json()["type"] == "token" # Antworttext kommt als token-Event(e)
def test_ws_voice_audio_stream_default_on(monkeypatch):
# Server-Default an -> Audio wird gestreamt, auch ohne expliziten audio_stream-Flag.
monkeypatch.setattr(settings, "audio_stream_default", True)
opts = _install_voice_stubs(monkeypatch)
with client.websocket_connect("/ws/voice") as ws:
ws.send_json({"type": "start", "format": "wav", **opts}) # kein audio_stream gesetzt
ws.send_bytes(b"PCMDATA")
ws.send_json({"type": "end"})
assert ws.receive_json()["type"] == "transcript"
assert ws.receive_json()["type"] == "ack"
assert ws.receive_json()["type"] == "audio" # Audio VOR semantic -> Streaming aktiv
def test_ws_voice_empty_buffer_errors(monkeypatch):
opts = _install_voice_stubs(monkeypatch)
with client.websocket_connect("/ws/voice") as ws:
ws.send_json({"type": "end", **opts}) # kein Audio gesendet
err = ws.receive_json()
assert err["type"] == "error" and "no audio" in err["detail"]