my_voice_assistant_v3/app/api/ws.py

308 lines
11 KiB
Python
Raw Normal View History

"""WebSocket-Echtzeit-Chat und -Sprache.
- /ws/chat : Text rein (JSON pro Turn), Antwort als Event-Folge zurueck.
- /ws/voice: Audio rein (binaere Chunks + Control), Transkription -> selbe Pipeline.
Antwort-Events: ack -> [token*] -> [audio*] -> semantic -> done.
Mit {"stream":true} kommen LLM-Token live, mit {"audio_stream":true} das Audio
satzweise (chunked TTS). /ws/voice sendet zuvor ein transcript-Event.
Barge-in: Ein {"type":"interrupt"}-Frame oder eine neue Eingabe bricht eine laufende
Antwort ab (-> interrupted-Event). Der Antwort-Turn laeuft als abbrechbarer Task.
Spaeter (eigene Increments): echte partielle Live-Transkripte (Streaming-STT-Dienst),
WebRTC.
"""
import asyncio
import json
from fastapi import APIRouter, WebSocket, WebSocketDisconnect
from app.config import settings
from app.errors import RoutingError
from app.dependencies import (
get_store,
resolve_route,
build_orchestrator,
resolve_output_endpoint,
)
from app.store import SessionOwnershipError
from app.audio.vad import EnergyVAD
from app.quota import enforce_quota, record_usage, QuotaExceededError
from app.safety.emergency import handle_emergency
router = APIRouter()
_OVERRIDE_KEYS = (
"input_endpoint",
"output_endpoint",
"language",
"stt_provider",
"llm_provider",
"tts_provider",
)
def _authenticate(token: str | None):
store = get_store()
if not settings.auth_enabled:
return store.ensure_anonymous_user()
if not token:
return None
return store.get_user_by_token(token)
async def _resolve(user, session_id, options):
"""Loest Route + Orchestrator + Output-Endpunkt auf (kann RoutingError/Ownership werfen)."""
overrides = {key: options.get(key) for key in _OVERRIDE_KEYS}
route = resolve_route(user, session_id, overrides)
orchestrator = build_orchestrator(route)
output = await resolve_output_endpoint(route)
return route, orchestrator, output
async def _run_turn(websocket, store, user, session_id, route, orchestrator, output, text, options):
"""Faehrt einen Antwort-Turn und streamt die Events an den Client."""
conversation = (
store.get_recent_messages(session_id, settings.history_max_messages)
if session_id
else []
)
memories = store.get_memories(user.id)
llm_context = list(conversation)
if memories:
memory_text = "Was du ueber den Nutzer weisst:\n" + "\n".join(
f"- {m.content}" for m in memories
)
llm_context = [{"role": "system", "content": memory_text}] + llm_context
# Notfall-Erkennung zuerst (immer eskalieren, auch bei Quota-Limit).
emergency = handle_emergency(user, text, store)
if emergency:
await websocket.send_json({"type": "emergency", "category": emergency["category"]})
else:
try:
enforce_quota(user, store)
except QuotaExceededError as exc:
await websocket.send_json({"type": "error", "status": 429, "detail": str(exc)})
return
await websocket.send_json({"type": "ack", "route": route.as_dict()})
voice = options.get("voice") or settings.openrouter_tts_voice
stream = bool(options.get("stream"))
audio_stream = bool(options.get("audio_stream"))
on_token = None
if stream:
async def on_token(delta):
await websocket.send_json({"type": "token", "text": delta})
on_audio = None
if audio_stream:
audio_seq = 0
async def on_audio(chunk):
nonlocal audio_seq
await websocket.send_json({"type": "audio", "seq": audio_seq})
audio_seq += 1
await websocket.send_bytes(chunk)
try:
if stream or audio_stream:
trace, audio = await orchestrator.chat_stream(
text,
language=route.language,
voice=voice,
output=output,
history=llm_context,
on_token=on_token,
on_audio=on_audio,
)
else:
trace, audio = await orchestrator.chat_text(
text,
language=route.language,
voice=voice,
output=output,
history=llm_context,
)
except Exception as exc:
await websocket.send_json({"type": "error", "status": 502, "detail": str(exc)})
return
record_usage(user, store, len(text) + len(trace.semantic_response or ""))
if session_id:
store.append_message(session_id, user.id, "user", text)
store.append_message(session_id, user.id, "assistant", trace.semantic_response)
await websocket.send_json(
{"type": "semantic", "text": trace.semantic_response, "spoken": trace.spoken_response}
)
if not audio_stream:
await websocket.send_bytes(audio)
await websocket.send_json({"type": "done", "audio_format": "pcm", "sample_rate": 24000})
async def _cancel_active(task, websocket) -> None:
"""Bricht einen laufenden Antwort-Turn ab (Barge-in) und meldet 'interrupted'."""
if task is None or task.done():
return
task.cancel()
try:
await task
except asyncio.CancelledError:
pass
await websocket.send_json({"type": "interrupted"})
async def _chat_turn(websocket, store, user, session_id, text, options):
try:
route, orchestrator, output = await _resolve(user, session_id, options)
except SessionOwnershipError as exc:
await websocket.send_json({"type": "error", "status": 403, "detail": str(exc)})
return
except RoutingError as exc:
await websocket.send_json({"type": "error", "status": 422, "detail": str(exc)})
return
await _run_turn(websocket, store, user, session_id, route, orchestrator, output, text, options)
async def _voice_turn(websocket, store, user, session_id, audio, fmt, options):
try:
route, orchestrator, output = await _resolve(user, session_id, options)
except SessionOwnershipError as exc:
await websocket.send_json({"type": "error", "status": 403, "detail": str(exc)})
return
except RoutingError as exc:
await websocket.send_json({"type": "error", "status": 422, "detail": str(exc)})
return
try:
transcript = await orchestrator.stt.transcribe(audio, fmt=fmt, language=route.language)
except Exception as exc:
await websocket.send_json({"type": "error", "status": 502, "detail": str(exc)})
return
await websocket.send_json({"type": "transcript", "text": transcript})
if not transcript or not transcript.strip():
await websocket.send_json({"type": "error", "detail": "empty transcript"})
return
await _run_turn(websocket, store, user, session_id, route, orchestrator, output, transcript, options)
@router.websocket("/ws/chat")
async def ws_chat(websocket: WebSocket, session_id: str | None = None, token: str | None = None):
user = _authenticate(token)
if user is None:
await websocket.close(code=1008)
return
await websocket.accept()
store = get_store()
active = None
try:
while True:
msg = await websocket.receive_json()
if msg.get("type") == "interrupt":
await _cancel_active(active, websocket)
active = None
continue
text = (msg.get("text") or "").strip()
if not text:
await websocket.send_json({"type": "error", "detail": "empty text"})
continue
await _cancel_active(active, websocket) # Barge-in bei neuer Eingabe
active = asyncio.create_task(
_chat_turn(websocket, store, user, session_id, text, msg)
)
except WebSocketDisconnect:
if active and not active.done():
active.cancel()
return
@router.websocket("/ws/voice")
async def ws_voice(websocket: WebSocket, session_id: str | None = None, token: str | None = None):
user = _authenticate(token)
if user is None:
await websocket.close(code=1008)
return
await websocket.accept()
store = get_store()
audio_buffer = bytearray()
fmt = "wav"
active = None
vad = None
start_options: dict = {} # Konfig aus dem start-Frame (Provider, audio_stream, language)
async def _start_voice(audio: bytes, options: dict):
nonlocal active
await _cancel_active(active, websocket) # Barge-in bei neuer Aeusserung
active = asyncio.create_task(
_voice_turn(websocket, store, user, session_id, audio, fmt, options)
)
try:
while True:
message = await websocket.receive()
if message["type"] == "websocket.disconnect":
if active and not active.done():
active.cancel()
return
if message.get("bytes") is not None:
audio_buffer.extend(message["bytes"])
# VAD: Aeusserungsende automatisch erkennen (opt-in via start-Frame).
if vad is not None and vad.feed(message["bytes"]):
audio = bytes(audio_buffer)
audio_buffer.clear()
vad.reset()
await _start_voice(audio, start_options)
continue
raw = message.get("text")
if raw is None:
continue
try:
control = json.loads(raw)
except ValueError:
await websocket.send_json({"type": "error", "detail": "invalid control frame"})
continue
ctype = control.get("type")
if ctype == "interrupt":
await _cancel_active(active, websocket)
active = None
continue
if ctype == "start":
audio_buffer.clear()
start_options = control # Provider/audio_stream/language fuer den Turn merken
fmt = control.get("format", "wav")
if control.get("vad"):
vad = EnergyVAD(
sample_rate=control.get("sample_rate", 16000),
threshold=control.get("vad_threshold", 500.0),
silence_ms=control.get("vad_silence_ms", 700.0),
)
else:
vad = None
continue
if ctype != "end":
continue
if not audio_buffer:
await websocket.send_json({"type": "error", "detail": "no audio received"})
continue
audio = bytes(audio_buffer)
audio_buffer.clear()
if vad is not None:
vad.reset()
# start-Frame-Konfig + end-Frame zusammenfuehren (end kann ueberschreiben)
await _start_voice(audio, {**start_options, **control})
except WebSocketDisconnect:
if active and not active.done():
active.cancel()
return