feat: Token-Level-LLM-Streaming über WebSocket (#4 Ausbau)
- LLMProvider.stream: Basis-Default (Fallback über complete) + SSE-Streaming
fuer OpenRouter und lokalen OpenAI-kompatiblen Provider; gemeinsamer Parser sse_delta
- Orchestrator.chat_stream: LLM-Token live via on_token-Callback, danach
Spoken-Adapter/Normalizer/TTS/Output; Fallback fuer Provider ohne stream
- WS /ws/chat: opt-in {"stream":true} -> ack -> token* -> semantic -> audio -> done
- Tests: 43 gruen (+5: SSE-Parsing, Default-Fallback, WS-Token-Flow)
- Doku aktualisiert; .gitignore: *.wav (generierte Audio-Ausgaben)
Co-Authored-By: Claude Opus 4.8 <noreply@anthropic.com>
This commit is contained in:
parent
531b57e08d
commit
379e002460
10 changed files with 302 additions and 24 deletions
3
.gitignore
vendored
3
.gitignore
vendored
|
|
@ -7,6 +7,9 @@ config/voice-assistant.toml
|
||||||
# Persistente Daten (SQLite-DB etc.)
|
# Persistente Daten (SQLite-DB etc.)
|
||||||
data/
|
data/
|
||||||
|
|
||||||
|
# Generierte Audio-Ausgaben (z. B. chat_client.py)
|
||||||
|
*.wav
|
||||||
|
|
||||||
# Python
|
# Python
|
||||||
__pycache__/
|
__pycache__/
|
||||||
*.py[cod]
|
*.py[cod]
|
||||||
|
|
|
||||||
|
|
@ -272,7 +272,9 @@ Diese Erinnerungen gibt der Assistent bei jedem Chat als Kontext mit — auch oh
|
||||||
|
|
||||||
**Echtzeit-Chat über WebSocket** (`/ws/chat`): dauerhafter Kanal, pro Nachricht
|
**Echtzeit-Chat über WebSocket** (`/ws/chat`): dauerhafter Kanal, pro Nachricht
|
||||||
`{"text": "..."}`; Antwort kommt als Event-Folge (`ack`, `semantic`, Audio, `done`).
|
`{"text": "..."}`; Antwort kommt als Event-Folge (`ack`, `semantic`, Audio, `done`).
|
||||||
Token per Query (`?token=…`), Gedächtnis per `?session_id=…`.
|
Token per Query (`?token=…`), Gedächtnis per `?session_id=…`. Mit
|
||||||
|
`{"text": "...", "stream": true}` kommt die Antwort schon während der Generierung
|
||||||
|
als `token`-Events (geringere wahrgenommene Latenz).
|
||||||
|
|
||||||
> **Für lokale Entwicklung** ist in der mitgelieferten `.env` `AUTH_ENABLED=false`
|
> **Für lokale Entwicklung** ist in der mitgelieferten `.env` `AUTH_ENABLED=false`
|
||||||
> gesetzt — dann ist kein Token nötig (anonymer Nutzer).
|
> gesetzt — dann ist kein Token nötig (anonymer Nutzer).
|
||||||
|
|
|
||||||
|
|
@ -186,7 +186,8 @@ Device Router (strikt, Singleton); Output-Lifecycle; **Authentifizierung
|
||||||
(Bearer-Token) + persistenter SQLite-Store für Nutzer/Sessions + Mandanten-Trennung
|
(Bearer-Token) + persistenter SQLite-Store für Nutzer/Sessions + Mandanten-Trennung
|
||||||
+ dauerhafte Nutzer-Präferenzen**; **Gesprächsgedächtnis pro Session (Verlauf im
|
+ dauerhafte Nutzer-Präferenzen**; **Gesprächsgedächtnis pro Session (Verlauf im
|
||||||
Store, fließt ins LLM)**; **Langzeit-Erinnerungen pro Nutzer (als LLM-Kontext)**;
|
Store, fließt ins LLM)**; **Langzeit-Erinnerungen pro Nutzer (als LLM-Kontext)**;
|
||||||
**WebSocket-Streaming-Chat (`/ws/chat`)**; automatisierte Tests.
|
**WebSocket-Streaming-Chat (`/ws/chat`) inkl. Token-Level-LLM-Streaming (SSE,
|
||||||
|
opt-in via `stream:true`)**; automatisierte Tests.
|
||||||
|
|
||||||
**Platzhalter (Gerüst):** Audio-Endpunkte (`local-default`, `bluetooth`,
|
**Platzhalter (Gerüst):** Audio-Endpunkte (`local-default`, `bluetooth`,
|
||||||
`mobile-ws`, `mobile-webrtc`) liefern leere Chunks — nur Auswahl/Lifecycle sind
|
`mobile-ws`, `mobile-webrtc`) liefern leere Chunks — nur Auswahl/Lifecycle sind
|
||||||
|
|
@ -201,7 +202,7 @@ Reihenfolge der Weiterentwicklung:
|
||||||
1. **(erledigt)** Konfig- & Routing-Fundament: Profile, Device Router, Registry, Pro-Request-Override.
|
1. **(erledigt)** Konfig- & Routing-Fundament: Profile, Device Router, Registry, Pro-Request-Override.
|
||||||
2. **(erledigt)** Cloud-Fundament: Bearer-Token-Auth, Mehrbenutzer, persistenter SQLite-Store, Mandanten-Trennung, dauerhafte Nutzer-Präferenzen. Offen: Skalierung auf gemeinsamen Store (Postgres/Redis) für mehrere Instanzen.
|
2. **(erledigt)** Cloud-Fundament: Bearer-Token-Auth, Mehrbenutzer, persistenter SQLite-Store, Mandanten-Trennung, dauerhafte Nutzer-Präferenzen. Offen: Skalierung auf gemeinsamen Store (Postgres/Redis) für mehrere Instanzen.
|
||||||
3. **(erledigt)** Konversationsgedächtnis: Kurzzeit-Gesprächsverlauf pro Session + Langzeit-Erinnerungen pro Nutzer (manuell gepflegt, als LLM-Kontext). Offen: **automatische** Extraktion/Zusammenfassung von Erinnerungen aus Gesprächen.
|
3. **(erledigt)** Konversationsgedächtnis: Kurzzeit-Gesprächsverlauf pro Session + Langzeit-Erinnerungen pro Nutzer (manuell gepflegt, als LLM-Kontext). Offen: **automatische** Extraktion/Zusammenfassung von Erinnerungen aus Gesprächen.
|
||||||
4. **(teilweise erledigt)** Echtzeit: WebSocket-Streaming-Chat (`/ws/chat`) mit Event-Folge (ack/semantic/audio/done) ist umgesetzt. Offen: **Token-Level-LLM-Streaming**, **Audio-Eingang/Streaming-STT**, **Barge-in/Turn-Manager**, **WebRTC**.
|
4. **(teilweise erledigt)** Echtzeit: WebSocket-Streaming-Chat (`/ws/chat`) mit Event-Folge (ack/semantic/audio/done) **und Token-Level-LLM-Streaming (SSE, opt-in `stream:true`)** sind umgesetzt. Offen: **Audio-Streaming (chunked TTS)**, **Audio-Eingang/Streaming-STT**, **Barge-in/Turn-Manager**, **WebRTC**.
|
||||||
5. **Resilienz:** Fallback-Policy (remote KI fällt aus → lokaler/alternativer Provider), Metriken/Tracing.
|
5. **Resilienz:** Fallback-Policy (remote KI fällt aus → lokaler/alternativer Provider), Metriken/Tracing.
|
||||||
6. **Betrieb:** Kosten-/Quota-Kontrolle pro Nutzer; Notfall-/Eskalationskonzept (Senioren-Kontext).
|
6. **Betrieb:** Kosten-/Quota-Kontrolle pro Nutzer; Notfall-/Eskalationskonzept (Senioren-Kontext).
|
||||||
7. **TransportRouter** als eigene lokal/remote-Achse aktivieren.
|
7. **TransportRouter** als eigene lokal/remote-Achse aktivieren.
|
||||||
|
|
|
||||||
10
README.md
10
README.md
|
|
@ -132,8 +132,14 @@ der Server streamt strukturierte Events zurück: `ack` → `semantic` → Audio
|
||||||
→ `done`. Auth (Token-Query `?token=…`), Session-Gedächtnis (`?session_id=…`) und
|
→ `done`. Auth (Token-Query `?token=…`), Session-Gedächtnis (`?session_id=…`) und
|
||||||
Erinnerungen gelten wie bei `POST /api/chat`.
|
Erinnerungen gelten wie bei `POST /api/chat`.
|
||||||
|
|
||||||
> Token-Level-LLM-Streaming, Audio-Eingang/Streaming-STT, Barge-in und WebRTC sind
|
**Token-Streaming:** Mit `{"text": "...", "stream": true}` schickt der Server die
|
||||||
> als nächste Increments vorgesehen (siehe Architektur-Dokument).
|
LLM-Antwort schon während der Generierung als `token`-Events
|
||||||
|
(`ack` → `token*` → `semantic` → Audio → `done`) — spürbar geringere wahrgenommene
|
||||||
|
Latenz. OpenRouter und der lokale OpenAI-kompatible Provider streamen via SSE;
|
||||||
|
Provider ohne Streaming liefern die komplette Antwort als ein `token`-Event.
|
||||||
|
|
||||||
|
> Audio-Streaming (chunked TTS), Audio-Eingang/Streaming-STT, Barge-in und WebRTC
|
||||||
|
> sind als nächste Increments vorgesehen (siehe Architektur-Dokument).
|
||||||
|
|
||||||
## Authentifizierung
|
## Authentifizierung
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -95,7 +95,21 @@ async def ws_chat(
|
||||||
await websocket.send_json({"type": "ack", "route": route.as_dict()})
|
await websocket.send_json({"type": "ack", "route": route.as_dict()})
|
||||||
|
|
||||||
voice = msg.get("voice") or settings.openrouter_tts_voice
|
voice = msg.get("voice") or settings.openrouter_tts_voice
|
||||||
|
stream = bool(msg.get("stream"))
|
||||||
try:
|
try:
|
||||||
|
if stream:
|
||||||
|
async def on_token(delta):
|
||||||
|
await websocket.send_json({"type": "token", "text": delta})
|
||||||
|
|
||||||
|
trace, audio = await orchestrator.chat_stream(
|
||||||
|
text,
|
||||||
|
language=route.language,
|
||||||
|
voice=voice,
|
||||||
|
output=output,
|
||||||
|
history=llm_context,
|
||||||
|
on_token=on_token,
|
||||||
|
)
|
||||||
|
else:
|
||||||
trace, audio = await orchestrator.chat_text(
|
trace, audio = await orchestrator.chat_text(
|
||||||
text,
|
text,
|
||||||
language=route.language,
|
language=route.language,
|
||||||
|
|
|
||||||
|
|
@ -105,3 +105,52 @@ class Orchestrator:
|
||||||
)
|
)
|
||||||
await self._emit_to_output(audio, output)
|
await self._emit_to_output(audio, output)
|
||||||
return trace, audio
|
return trace, audio
|
||||||
|
|
||||||
|
async def chat_stream(
|
||||||
|
self,
|
||||||
|
text: str,
|
||||||
|
language: str | None = None,
|
||||||
|
voice: str | None = None,
|
||||||
|
output=None,
|
||||||
|
history: list[dict] | None = None,
|
||||||
|
on_token=None,
|
||||||
|
):
|
||||||
|
"""Wie chat_text, aber die LLM-Antwort wird tokenweise gestreamt.
|
||||||
|
|
||||||
|
`on_token(delta)` (async) wird pro Token-Delta aufgerufen. Audio/Output
|
||||||
|
werden erst nach der vollstaendigen Antwort erzeugt (TTS ist nicht streamend).
|
||||||
|
"""
|
||||||
|
trace = PipelineTrace()
|
||||||
|
trace.raw_transcript = text
|
||||||
|
trace.cleaned_transcript = await self.input_cleaner.run(text or "")
|
||||||
|
|
||||||
|
parts: list[str] = []
|
||||||
|
stream_fn = getattr(self.llm, "stream", None)
|
||||||
|
if stream_fn is not None:
|
||||||
|
async for delta in stream_fn(trace.cleaned_transcript or "", history=history):
|
||||||
|
parts.append(delta)
|
||||||
|
if on_token:
|
||||||
|
await on_token(delta)
|
||||||
|
else:
|
||||||
|
# Provider ohne Streaming -> komplette Antwort als ein Token.
|
||||||
|
result = await self.llm.complete(trace.cleaned_transcript or "", history=history)
|
||||||
|
parts.append(result)
|
||||||
|
if on_token:
|
||||||
|
await on_token(result)
|
||||||
|
|
||||||
|
trace.semantic_response = "".join(parts)
|
||||||
|
if not trace.semantic_response:
|
||||||
|
raise RuntimeError("LLM returned an empty response")
|
||||||
|
|
||||||
|
trace.spoken_response = await self.spoken_adapter.run(
|
||||||
|
trace.semantic_response,
|
||||||
|
language=language,
|
||||||
|
)
|
||||||
|
trace.tts_ready_text = await self.tts_normalizer.run(
|
||||||
|
trace.spoken_response,
|
||||||
|
language=language,
|
||||||
|
)
|
||||||
|
|
||||||
|
audio = await self.tts.synthesize(trace.tts_ready_text, voice=voice)
|
||||||
|
await self._emit_to_output(audio, output)
|
||||||
|
return trace, audio
|
||||||
|
|
|
||||||
|
|
@ -1,4 +1,21 @@
|
||||||
|
import json
|
||||||
from abc import ABC, abstractmethod
|
from abc import ABC, abstractmethod
|
||||||
|
from collections.abc import AsyncIterator
|
||||||
|
|
||||||
|
|
||||||
|
def sse_delta(line: str) -> str | None:
|
||||||
|
"""Extrahiert das Token-Delta aus einer OpenAI-kompatiblen SSE-Zeile (oder None)."""
|
||||||
|
if not line.startswith("data:"):
|
||||||
|
return None
|
||||||
|
data = line[len("data:"):].strip()
|
||||||
|
if not data or data == "[DONE]":
|
||||||
|
return None
|
||||||
|
try:
|
||||||
|
obj = json.loads(data)
|
||||||
|
return obj["choices"][0]["delta"].get("content")
|
||||||
|
except (ValueError, KeyError, IndexError, TypeError):
|
||||||
|
return None
|
||||||
|
|
||||||
|
|
||||||
class LLMProvider(ABC):
|
class LLMProvider(ABC):
|
||||||
@abstractmethod
|
@abstractmethod
|
||||||
|
|
@ -8,3 +25,15 @@ class LLMProvider(ABC):
|
||||||
history: list[dict] | None = None,
|
history: list[dict] | None = None,
|
||||||
session_id: str | None = None,
|
session_id: str | None = None,
|
||||||
) -> str: ...
|
) -> str: ...
|
||||||
|
|
||||||
|
async def stream(
|
||||||
|
self,
|
||||||
|
text: str,
|
||||||
|
history: list[dict] | None = None,
|
||||||
|
session_id: str | None = None,
|
||||||
|
) -> AsyncIterator[str]:
|
||||||
|
"""Token-Stream. Default: kein echtes Streaming -> komplette Antwort als ein Chunk.
|
||||||
|
|
||||||
|
Provider mit SSE-Unterstuetzung ueberschreiben diese Methode.
|
||||||
|
"""
|
||||||
|
yield await self.complete(text, history=history, session_id=session_id)
|
||||||
|
|
|
||||||
|
|
@ -1,5 +1,9 @@
|
||||||
|
from collections.abc import AsyncIterator
|
||||||
|
|
||||||
import httpx
|
import httpx
|
||||||
from app.providers.llm.base import LLMProvider
|
|
||||||
|
from app.providers.llm.base import LLMProvider, sse_delta
|
||||||
|
|
||||||
|
|
||||||
class LocalOpenAICompatibleLLM(LLMProvider):
|
class LocalOpenAICompatibleLLM(LLMProvider):
|
||||||
def __init__(self, base_url: str, api_key: str, model: str):
|
def __init__(self, base_url: str, api_key: str, model: str):
|
||||||
|
|
@ -7,17 +11,20 @@ class LocalOpenAICompatibleLLM(LLMProvider):
|
||||||
self.api_key = api_key
|
self.api_key = api_key
|
||||||
self.model = model
|
self.model = model
|
||||||
|
|
||||||
|
def _build_messages(self, text: str, history: list[dict] | None) -> list[dict]:
|
||||||
|
messages = list(history) if history else []
|
||||||
|
messages.append({"role": "user", "content": text})
|
||||||
|
return messages
|
||||||
|
|
||||||
async def complete(
|
async def complete(
|
||||||
self,
|
self,
|
||||||
text: str,
|
text: str,
|
||||||
history: list[dict] | None = None,
|
history: list[dict] | None = None,
|
||||||
session_id: str | None = None,
|
session_id: str | None = None,
|
||||||
) -> str:
|
) -> str:
|
||||||
messages = list(history) if history else []
|
|
||||||
messages.append({"role": "user", "content": text})
|
|
||||||
payload = {
|
payload = {
|
||||||
"model": self.model,
|
"model": self.model,
|
||||||
"messages": messages,
|
"messages": self._build_messages(text, history),
|
||||||
"temperature": 0.3,
|
"temperature": 0.3,
|
||||||
}
|
}
|
||||||
async with httpx.AsyncClient(timeout=120) as client:
|
async with httpx.AsyncClient(timeout=120) as client:
|
||||||
|
|
@ -29,3 +36,33 @@ class LocalOpenAICompatibleLLM(LLMProvider):
|
||||||
response.raise_for_status()
|
response.raise_for_status()
|
||||||
data = response.json()
|
data = response.json()
|
||||||
return data["choices"][0]["message"]["content"]
|
return data["choices"][0]["message"]["content"]
|
||||||
|
|
||||||
|
async def stream(
|
||||||
|
self,
|
||||||
|
text: str,
|
||||||
|
history: list[dict] | None = None,
|
||||||
|
session_id: str | None = None,
|
||||||
|
) -> AsyncIterator[str]:
|
||||||
|
payload = {
|
||||||
|
"model": self.model,
|
||||||
|
"messages": self._build_messages(text, history),
|
||||||
|
"temperature": 0.3,
|
||||||
|
"stream": True,
|
||||||
|
}
|
||||||
|
async with httpx.AsyncClient(timeout=120) as client:
|
||||||
|
async with client.stream(
|
||||||
|
"POST",
|
||||||
|
f"{self.base_url}/chat/completions",
|
||||||
|
headers={"Authorization": f"Bearer {self.api_key}"},
|
||||||
|
json=payload,
|
||||||
|
) as response:
|
||||||
|
if response.status_code >= 400:
|
||||||
|
body = await response.aread()
|
||||||
|
raise RuntimeError(
|
||||||
|
f"Local LLM error {response.status_code}: "
|
||||||
|
f"{body.decode(errors='replace')}"
|
||||||
|
)
|
||||||
|
async for line in response.aiter_lines():
|
||||||
|
delta = sse_delta(line)
|
||||||
|
if delta:
|
||||||
|
yield delta
|
||||||
|
|
|
||||||
|
|
@ -1,6 +1,8 @@
|
||||||
|
from collections.abc import AsyncIterator
|
||||||
|
|
||||||
import httpx
|
import httpx
|
||||||
|
|
||||||
from app.providers.llm.base import LLMProvider
|
from app.providers.llm.base import LLMProvider, sse_delta
|
||||||
|
|
||||||
|
|
||||||
SYSTEM_PROMPT = """
|
SYSTEM_PROMPT = """
|
||||||
|
|
@ -49,12 +51,7 @@ class OpenRouterLLMProvider(LLMProvider):
|
||||||
self.api_key = (api_key or "").strip()
|
self.api_key = (api_key or "").strip()
|
||||||
self.model = (model or "").strip()
|
self.model = (model or "").strip()
|
||||||
|
|
||||||
async def complete(
|
def _build_messages(self, text: str, history: list[dict] | None) -> list[dict]:
|
||||||
self,
|
|
||||||
text: str,
|
|
||||||
history: list[dict] | None = None,
|
|
||||||
session_id: str | None = None,
|
|
||||||
) -> str:
|
|
||||||
if not self.api_key:
|
if not self.api_key:
|
||||||
raise ValueError("OPENROUTER_API_KEY is empty")
|
raise ValueError("OPENROUTER_API_KEY is empty")
|
||||||
if not self.model:
|
if not self.model:
|
||||||
|
|
@ -66,10 +63,17 @@ class OpenRouterLLMProvider(LLMProvider):
|
||||||
if history:
|
if history:
|
||||||
messages.extend(history)
|
messages.extend(history)
|
||||||
messages.append({"role": "user", "content": text.strip()})
|
messages.append({"role": "user", "content": text.strip()})
|
||||||
|
return messages
|
||||||
|
|
||||||
|
async def complete(
|
||||||
|
self,
|
||||||
|
text: str,
|
||||||
|
history: list[dict] | None = None,
|
||||||
|
session_id: str | None = None,
|
||||||
|
) -> str:
|
||||||
payload = {
|
payload = {
|
||||||
"model": self.model,
|
"model": self.model,
|
||||||
"messages": messages,
|
"messages": self._build_messages(text, history),
|
||||||
}
|
}
|
||||||
|
|
||||||
timeout = httpx.Timeout(connect=10.0, read=120.0, write=30.0, pool=10.0)
|
timeout = httpx.Timeout(connect=10.0, read=120.0, write=30.0, pool=10.0)
|
||||||
|
|
@ -106,3 +110,42 @@ class OpenRouterLLMProvider(LLMProvider):
|
||||||
|
|
||||||
return str(content).strip()
|
return str(content).strip()
|
||||||
|
|
||||||
|
async def stream(
|
||||||
|
self,
|
||||||
|
text: str,
|
||||||
|
history: list[dict] | None = None,
|
||||||
|
session_id: str | None = None,
|
||||||
|
) -> AsyncIterator[str]:
|
||||||
|
payload = {
|
||||||
|
"model": self.model,
|
||||||
|
"messages": self._build_messages(text, history),
|
||||||
|
"stream": True,
|
||||||
|
}
|
||||||
|
timeout = httpx.Timeout(connect=10.0, read=120.0, write=30.0, pool=10.0)
|
||||||
|
|
||||||
|
async with httpx.AsyncClient(timeout=timeout) as client:
|
||||||
|
try:
|
||||||
|
async with client.stream(
|
||||||
|
"POST",
|
||||||
|
"https://openrouter.ai/api/v1/chat/completions",
|
||||||
|
headers={
|
||||||
|
"Authorization": f"Bearer {self.api_key}",
|
||||||
|
"Content-Type": "application/json",
|
||||||
|
},
|
||||||
|
json=payload,
|
||||||
|
) as response:
|
||||||
|
if response.status_code >= 400:
|
||||||
|
body = await response.aread()
|
||||||
|
raise RuntimeError(
|
||||||
|
f"OpenRouter LLM error {response.status_code}: "
|
||||||
|
f"{body.decode(errors='replace')}"
|
||||||
|
)
|
||||||
|
async for line in response.aiter_lines():
|
||||||
|
delta = sse_delta(line)
|
||||||
|
if delta:
|
||||||
|
yield delta
|
||||||
|
except httpx.TimeoutException as exc:
|
||||||
|
raise RuntimeError("OpenRouter LLM timeout") from exc
|
||||||
|
except httpx.HTTPError as exc:
|
||||||
|
raise RuntimeError(f"OpenRouter LLM transport error: {exc}") from exc
|
||||||
|
|
||||||
|
|
|
||||||
94
tests/test_streaming.py
Normal file
94
tests/test_streaming.py
Normal file
|
|
@ -0,0 +1,94 @@
|
||||||
|
import asyncio
|
||||||
|
|
||||||
|
from fastapi.testclient import TestClient
|
||||||
|
|
||||||
|
import app.dependencies as deps
|
||||||
|
from app.main import app
|
||||||
|
from app.providers.llm.base import sse_delta, LLMProvider
|
||||||
|
|
||||||
|
client = TestClient(app)
|
||||||
|
|
||||||
|
|
||||||
|
def test_sse_delta_parsing():
|
||||||
|
assert sse_delta('data: {"choices":[{"delta":{"content":"Hal"}}]}') == "Hal"
|
||||||
|
assert sse_delta("data: [DONE]") is None
|
||||||
|
assert sse_delta("") is None
|
||||||
|
assert sse_delta(": keep-alive") is None
|
||||||
|
assert sse_delta('data: {"choices":[{"delta":{}}]}') is None
|
||||||
|
assert sse_delta("data: nicht-json") is None
|
||||||
|
|
||||||
|
|
||||||
|
def test_base_stream_default_yields_full_completion():
|
||||||
|
class P(LLMProvider):
|
||||||
|
async def complete(self, text, history=None, session_id=None):
|
||||||
|
return "ganze Antwort"
|
||||||
|
|
||||||
|
async def run():
|
||||||
|
return [delta async for delta in P().stream("x")]
|
||||||
|
|
||||||
|
assert asyncio.run(run()) == ["ganze Antwort"]
|
||||||
|
|
||||||
|
|
||||||
|
def _install_streaming(monkeypatch, tokens):
|
||||||
|
class StreamLLM:
|
||||||
|
async def complete(self, text, history=None, session_id=None):
|
||||||
|
return "".join(tokens)
|
||||||
|
|
||||||
|
async def stream(self, text, history=None, session_id=None):
|
||||||
|
for tok in tokens:
|
||||||
|
yield tok
|
||||||
|
|
||||||
|
class StubTTS:
|
||||||
|
async def synthesize(self, text, voice=None, audio_format="pcm"):
|
||||||
|
return b"AUD"
|
||||||
|
|
||||||
|
monkeypatch.setitem(deps.LLM_REGISTRY, "stream", lambda s: StreamLLM())
|
||||||
|
monkeypatch.setitem(deps.TTS_REGISTRY, "stub", lambda s: StubTTS())
|
||||||
|
return {"llm_provider": "stream", "tts_provider": "stub", "output_endpoint": "loopback"}
|
||||||
|
|
||||||
|
|
||||||
|
def test_ws_stream_emits_token_events(monkeypatch):
|
||||||
|
base = _install_streaming(monkeypatch, ["Gu", "ten ", "Tag"])
|
||||||
|
with client.websocket_connect("/ws/chat") as ws:
|
||||||
|
ws.send_json({"text": "Hallo", "stream": True, **base})
|
||||||
|
assert ws.receive_json()["type"] == "ack"
|
||||||
|
|
||||||
|
tokens = []
|
||||||
|
event = ws.receive_json()
|
||||||
|
while event["type"] == "token":
|
||||||
|
tokens.append(event["text"])
|
||||||
|
event = ws.receive_json()
|
||||||
|
|
||||||
|
assert tokens == ["Gu", "ten ", "Tag"]
|
||||||
|
assert event["type"] == "semantic" and event["text"] == "Guten Tag"
|
||||||
|
assert ws.receive_bytes() == b"AUD"
|
||||||
|
assert ws.receive_json()["type"] == "done"
|
||||||
|
|
||||||
|
|
||||||
|
def test_ws_without_stream_flag_has_no_tokens(monkeypatch):
|
||||||
|
base = _install_streaming(monkeypatch, ["a", "b"])
|
||||||
|
with client.websocket_connect("/ws/chat") as ws:
|
||||||
|
ws.send_json({"text": "Hallo", **base}) # kein stream-Flag
|
||||||
|
assert ws.receive_json()["type"] == "ack"
|
||||||
|
assert ws.receive_json()["type"] == "semantic" # direkt, keine token-Events
|
||||||
|
|
||||||
|
|
||||||
|
def test_ws_stream_fallback_for_nonstreaming_llm(monkeypatch):
|
||||||
|
class OnlyComplete:
|
||||||
|
async def complete(self, text, history=None, session_id=None):
|
||||||
|
return "komplett"
|
||||||
|
|
||||||
|
class StubTTS:
|
||||||
|
async def synthesize(self, text, voice=None, audio_format="pcm"):
|
||||||
|
return b"X"
|
||||||
|
|
||||||
|
monkeypatch.setitem(deps.LLM_REGISTRY, "oc", lambda s: OnlyComplete())
|
||||||
|
monkeypatch.setitem(deps.TTS_REGISTRY, "stub", lambda s: StubTTS())
|
||||||
|
base = {"llm_provider": "oc", "tts_provider": "stub", "output_endpoint": "loopback"}
|
||||||
|
|
||||||
|
with client.websocket_connect("/ws/chat") as ws:
|
||||||
|
ws.send_json({"text": "x", "stream": True, **base})
|
||||||
|
assert ws.receive_json()["type"] == "ack"
|
||||||
|
token = ws.receive_json()
|
||||||
|
assert token["type"] == "token" and token["text"] == "komplett"
|
||||||
|
assert ws.receive_json()["type"] == "semantic"
|
||||||
Loading…
Add table
Add a link
Reference in a new issue