feat: stream provider responses
This commit is contained in:
+17
-16
@@ -106,6 +106,9 @@ Rootserver:
|
||||
Behandlungs- und medizinische Notfallfragen sind zwingend lokal `never`.
|
||||
- `/privacy` trennt Inhaltsklassifikation, tatsächlichen Provider, Unterdrückung
|
||||
durch Modus `local` und technischen Provider-Fallback.
|
||||
- Ollama, Gemini, Hybridrouter, Chatservice und CLI streamen sichtbare Antwortteile.
|
||||
Verdeckte Gedanken werden nicht angefordert oder ausgegeben; Terminal-Markdown
|
||||
wird lesbar bereinigt. Strg+C verwirft Teilantworten atomar und kehrt zum Chat zurück.
|
||||
|
||||
## Aktuelle Architektur
|
||||
|
||||
@@ -114,6 +117,8 @@ Rootserver:
|
||||
- Modell-Provider als kleine austauschbare Schnittstelle
|
||||
- Ollama-Provider akzeptiert nur lokale HTTP-Loopback-Adressen
|
||||
- SQLite-Sitzungsspeicher mit atomaren Benutzer-/Assistentenpaaren
|
||||
- Streaming wird erst nach vollständigem Abschluss atomar gespeichert; bei Abbruch
|
||||
bleibt weder die Benutzerfrage noch eine unvollständige Antwort im Verlauf.
|
||||
- Laufzeitdaten außerhalb von Git über `JAVIS_DATA_DIR` oder sicheren Plattformstandard
|
||||
- Hybridprovider bevorzugt Gemini für cloudgeeignete Inhalte und verwendet
|
||||
Ollama bei Datenschutz, Ablehnung, Offline-, Quota- und Providerfehlern
|
||||
@@ -164,6 +169,10 @@ Rootserver:
|
||||
- Lesender Ollama-Status prüft nur Loopback-Erreichbarkeit und Modellliste.
|
||||
- Sicheres PowerShell-Startskript und kompakte Startdokumentation ergänzt; keine
|
||||
PATH-, Registry-, Autostart-, Dienst- oder Richtlinienänderung.
|
||||
- Medizinische Inhalte vor Cloudclient-Erstellung strikt auf `never` gesetzt.
|
||||
- Echtes Ollama-/Gemini-Streaming, providerübergreifendes Fallback vor der ersten
|
||||
Ausgabe und sicheren Streaming-Abbruch nach Teilausgabe implementiert.
|
||||
- CLI-Streaming mit lesbarer Markdown-Bereinigung und sicherem Strg+C ergänzt.
|
||||
|
||||
## Aktuelle Tests
|
||||
|
||||
@@ -178,7 +187,8 @@ Letzter bestätigter Projektstand:
|
||||
- uv-Lock und `uv sync --dev`: bestanden
|
||||
- Python in `.venv`: 3.12.13
|
||||
- Ruff in `.venv`: 0.16.0
|
||||
- Unit-Tests: 69 bestanden; medizinisches `never` und Cloudclient-Sperre abgedeckt
|
||||
- Unit-Tests: 82 bestanden; medizinisches `never`, Streaming, Gedankenfilter,
|
||||
Cloudclient-Sperre, Fallback und atomarer Strg+C-Abbruch abgedeckt
|
||||
- PowerShell-Syntax des Startskripts: erfolgreich geparst
|
||||
- Ruff Lint: bestanden
|
||||
- Ruff Formatprüfung: bestanden
|
||||
@@ -220,7 +230,9 @@ Abnahmestatus:
|
||||
|
||||
## Git-Stand
|
||||
|
||||
- aktueller und stabiler Branch: `main`
|
||||
- aktueller Arbeitsbranch: `feat/chat-comfort`
|
||||
- stabiler Ausgangsstand: `main` bei `77b510b`
|
||||
- medizinischer Datenschutz-Fix: `9148193`
|
||||
- Feature-Branch ist als `7a14fe8` zu
|
||||
`origin/feat/gemini-privacy-router` gepusht
|
||||
- konfliktfreier Merge nach `main`: `564a3fd`
|
||||
@@ -249,23 +261,11 @@ Abnahmestatus:
|
||||
- `origin` verwendet HTTPS
|
||||
- kein Force-Push und keine umgeschriebene Historie
|
||||
|
||||
## Kompakter Changelog
|
||||
|
||||
- `49a1dd3`: erster Gitea-Verbindungstest
|
||||
- `6d04171`: sicheres Javis-Grundgerüst
|
||||
- `0e82748`: Werkzeug-, Token- und Handoff-Regeln
|
||||
- `e6eb995`: einmalige Archivierung und Verdichtung des Handoffs
|
||||
- aktuell: isolierte Python-Toolchain und reproduzierbare Entwicklungsumgebung
|
||||
- aktuell: portable Ollama-Laufzeit und `qwen3:8b` außerhalb von Git verifiziert
|
||||
- aktuell: lokaler CLI-Chat und persistente SQLite-Sitzungen vollständig abgenommen
|
||||
- aktuell: Gemini-Free-Hybridrouting, Datenschutz und Nullkostenmodus
|
||||
automatisiert und durch Pascal manuell vollständig abgenommen
|
||||
|
||||
## Offene Entscheidungen und Fehler
|
||||
|
||||
- endgültiger Produkt-/Repositoryname bleibt offen
|
||||
- Startskript startet Ollama bei Bedarf pro Prozess; kein Autostart oder Dienst
|
||||
- Antworten werden noch nicht gestreamt; Sitzungen besitzen noch keine Titel oder Suche
|
||||
- Antworten werden gestreamt; Sitzungen besitzen noch keine Titel oder Suche
|
||||
- genauer späterer Obsidian-Schreibbereich ist nicht freigegeben
|
||||
- normaler Secret-Provider ist festgelegt: Betriebssystem-Keyring; eine
|
||||
Klartext-XML wird nicht für API-Schlüssel verwendet
|
||||
@@ -289,7 +289,8 @@ Abnahmestatus:
|
||||
|
||||
## Nächster sinnvoller Auftrag
|
||||
|
||||
Auf `feat/chat-comfort` als nächstes Streaming mit sicherem Abbruch umsetzen.
|
||||
Auf `feat/chat-comfort` als Nächstes SQLite-Titelmigration, lokale automatische
|
||||
Titel sowie `/rename`, `/sessions`, `/load` und `/search` umsetzen.
|
||||
Medizinische Antwortqualität bleibt ein späteres Sicherheits-/Systemprompt-Thema.
|
||||
Noch keine Obsidian-Integration oder Tools beginnen.
|
||||
|
||||
|
||||
@@ -2,10 +2,15 @@
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Iterator
|
||||
from dataclasses import dataclass
|
||||
|
||||
from javis.memory.sqlite_store import ChatSession, SQLiteSessionStore
|
||||
from javis.providers.base import ChatMessage, LocalModelProvider
|
||||
from javis.providers.base import (
|
||||
ChatMessage,
|
||||
InvalidProviderResponseError,
|
||||
LocalModelProvider,
|
||||
)
|
||||
|
||||
|
||||
class SessionProviderMismatchError(RuntimeError):
|
||||
@@ -59,5 +64,43 @@ class ChatService:
|
||||
self.store.append_exchange(session_id, normalized, response)
|
||||
return response
|
||||
|
||||
def stream_send(self, session_id: str, text: str) -> Iterator[str]:
|
||||
normalized = text.strip()
|
||||
if not normalized:
|
||||
raise ValueError("Eine leere Nachricht wird nicht gesendet.")
|
||||
loaded = self.load_session(session_id)
|
||||
messages = [*loaded.messages, ChatMessage("user", normalized)]
|
||||
stream_method = getattr(self.provider, "stream_chat", None)
|
||||
stream = (
|
||||
stream_method(messages)
|
||||
if callable(stream_method)
|
||||
else iter((self.provider.chat(messages),))
|
||||
)
|
||||
chunks: list[str] = []
|
||||
completed = False
|
||||
try:
|
||||
for chunk in stream:
|
||||
if not isinstance(chunk, str):
|
||||
raise InvalidProviderResponseError(
|
||||
"Der Provider lieferte einen ungültigen Streaming-Abschnitt."
|
||||
)
|
||||
if not chunk:
|
||||
continue
|
||||
chunks.append(chunk)
|
||||
yield chunk
|
||||
completed = True
|
||||
finally:
|
||||
if not completed:
|
||||
close_stream = getattr(stream, "close", None)
|
||||
if callable(close_stream):
|
||||
close_stream()
|
||||
|
||||
response = "".join(chunks).strip()
|
||||
if not response:
|
||||
raise InvalidProviderResponseError(
|
||||
"Der Provider lieferte keine verwendbare Streaming-Antwort."
|
||||
)
|
||||
self.store.append_exchange(session_id, normalized, response)
|
||||
|
||||
def clear_session(self, session_id: str) -> None:
|
||||
self.store.clear_messages(session_id)
|
||||
|
||||
@@ -2,17 +2,19 @@
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Callable
|
||||
from collections.abc import Callable, Iterator
|
||||
from dataclasses import dataclass
|
||||
|
||||
from javis.memory.usage_store import ProviderEvent, SQLiteUsageStore, UsageStoreError
|
||||
from javis.providers.base import (
|
||||
ChatMessage,
|
||||
InvalidApiKeyError,
|
||||
InvalidProviderResponseError,
|
||||
LocalModelProvider,
|
||||
MissingApiKeyError,
|
||||
ProviderError,
|
||||
ProviderUsage,
|
||||
ResponseAbortedError,
|
||||
)
|
||||
from javis.security.privacy import CloudPolicy, PrivacyDecision, PrivacyRouter
|
||||
|
||||
@@ -165,6 +167,245 @@ class HybridProvider:
|
||||
)
|
||||
return response
|
||||
|
||||
def stream_chat(self, messages: list[ChatMessage]) -> Iterator[str]:
|
||||
current_text = messages[-1].content if messages else ""
|
||||
decision = self._privacy_router.classify(
|
||||
current_text,
|
||||
cloud_requested=self.mode == "gemini",
|
||||
)
|
||||
|
||||
if self.mode == "local":
|
||||
yield from self._stream_local(
|
||||
messages,
|
||||
decision,
|
||||
"lokaler Modus",
|
||||
fallback=False,
|
||||
cloud_suppressed_by_local_mode=True,
|
||||
)
|
||||
return
|
||||
if decision.policy is CloudPolicy.NEVER:
|
||||
yield from self._stream_local(
|
||||
messages,
|
||||
decision,
|
||||
decision.reason,
|
||||
fallback=True,
|
||||
)
|
||||
return
|
||||
|
||||
approved = decision.policy is CloudPolicy.ALLOWED
|
||||
if decision.policy is CloudPolicy.ASK:
|
||||
approved = self._approval_callback(decision)
|
||||
if not approved:
|
||||
yield from self._stream_local(
|
||||
messages,
|
||||
decision,
|
||||
"Cloudfreigabe abgelehnt",
|
||||
fallback=True,
|
||||
)
|
||||
return
|
||||
|
||||
unavailable_reason = self._cloud_unavailable_reason()
|
||||
if unavailable_reason:
|
||||
yield from self._stream_local(
|
||||
messages,
|
||||
decision,
|
||||
unavailable_reason,
|
||||
fallback=True,
|
||||
)
|
||||
return
|
||||
if len(current_text) > self._max_cloud_input_chars:
|
||||
yield from self._stream_local(
|
||||
messages,
|
||||
decision,
|
||||
"Nachricht überschreitet das lokale Cloud-Größenlimit",
|
||||
fallback=True,
|
||||
)
|
||||
return
|
||||
|
||||
cloud_messages = self._privacy_router.minimal_context(
|
||||
messages,
|
||||
current_policy=decision.policy,
|
||||
approved=approved,
|
||||
max_chars=self._max_cloud_input_chars,
|
||||
max_messages=self._max_cloud_context_messages,
|
||||
)
|
||||
if not cloud_messages:
|
||||
yield from self._stream_local(
|
||||
messages,
|
||||
decision,
|
||||
"kein freigegebener Cloudkontext",
|
||||
fallback=True,
|
||||
)
|
||||
return
|
||||
|
||||
emitted = False
|
||||
try:
|
||||
cloud_provider = self._cloud_provider_factory()
|
||||
for chunk in self._provider_chunks(cloud_provider, cloud_messages):
|
||||
emitted = True
|
||||
yield chunk
|
||||
except InvalidApiKeyError as exc:
|
||||
self._cloud_session_disabled = True
|
||||
if emitted:
|
||||
self._raise_interrupted_cloud(decision, exc)
|
||||
yield from self._stream_cloud_failure(
|
||||
messages,
|
||||
decision,
|
||||
exc,
|
||||
disable_notice=True,
|
||||
)
|
||||
return
|
||||
except (MissingApiKeyError, ProviderError) as exc:
|
||||
if emitted:
|
||||
self._raise_interrupted_cloud(decision, exc)
|
||||
yield from self._stream_cloud_failure(messages, decision, exc)
|
||||
return
|
||||
|
||||
if not emitted:
|
||||
error = InvalidProviderResponseError(
|
||||
"Gemini hat keine verwendbare Streaming-Antwort geliefert."
|
||||
)
|
||||
yield from self._stream_cloud_failure(messages, decision, error)
|
||||
return
|
||||
|
||||
usage = getattr(cloud_provider, "last_usage", ProviderUsage())
|
||||
self._record(
|
||||
ProviderEvent(
|
||||
provider=cloud_provider.name,
|
||||
model=cloud_provider.model,
|
||||
success=True,
|
||||
error_category=None,
|
||||
fallback=False,
|
||||
privacy_policy=decision.policy.value,
|
||||
input_tokens=usage.input_tokens,
|
||||
output_tokens=usage.output_tokens,
|
||||
)
|
||||
)
|
||||
self.last_route = RouteStatus(
|
||||
provider=cloud_provider.name,
|
||||
privacy_policy=decision.policy,
|
||||
fallback=False,
|
||||
cloud_suppressed_by_local_mode=False,
|
||||
technical_fallback=False,
|
||||
reason="Cloudstream erfolgreich",
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _provider_chunks(
|
||||
provider: LocalModelProvider,
|
||||
messages: list[ChatMessage],
|
||||
) -> Iterator[str]:
|
||||
stream_method = getattr(provider, "stream_chat", None)
|
||||
chunks = (
|
||||
stream_method(messages) if callable(stream_method) else iter((provider.chat(messages),))
|
||||
)
|
||||
completed = False
|
||||
try:
|
||||
for chunk in chunks:
|
||||
if not isinstance(chunk, str):
|
||||
raise InvalidProviderResponseError(
|
||||
"Der Provider lieferte einen ungültigen Streaming-Abschnitt."
|
||||
)
|
||||
if chunk:
|
||||
yield chunk
|
||||
completed = True
|
||||
finally:
|
||||
if not completed:
|
||||
close = getattr(chunks, "close", None)
|
||||
if callable(close):
|
||||
close()
|
||||
|
||||
def _stream_cloud_failure(
|
||||
self,
|
||||
messages: list[ChatMessage],
|
||||
decision: PrivacyDecision,
|
||||
error: ProviderError,
|
||||
*,
|
||||
disable_notice: bool = False,
|
||||
) -> Iterator[str]:
|
||||
self._record_cloud_error(decision, error)
|
||||
reason = "Gemini nicht verfügbar"
|
||||
if disable_notice:
|
||||
reason = "Gemini-Schlüssel abgelehnt; Cloud für diese Sitzung deaktiviert"
|
||||
yield from self._stream_local(
|
||||
messages,
|
||||
decision,
|
||||
reason,
|
||||
fallback=True,
|
||||
technical_fallback=True,
|
||||
)
|
||||
|
||||
def _raise_interrupted_cloud(
|
||||
self,
|
||||
decision: PrivacyDecision,
|
||||
error: ProviderError,
|
||||
) -> None:
|
||||
self._record_cloud_error(decision, error)
|
||||
self.last_route = RouteStatus(
|
||||
provider="gemini",
|
||||
privacy_policy=decision.policy,
|
||||
fallback=False,
|
||||
cloud_suppressed_by_local_mode=False,
|
||||
technical_fallback=True,
|
||||
reason="Cloudstream nach Teilausgabe abgebrochen",
|
||||
)
|
||||
raise ResponseAbortedError(
|
||||
"Der Gemini-Stream wurde nach einer Teilausgabe abgebrochen; "
|
||||
"es wurde kein lokaler Ersatz angehängt."
|
||||
) from error
|
||||
|
||||
def _stream_local(
|
||||
self,
|
||||
messages: list[ChatMessage],
|
||||
decision: PrivacyDecision,
|
||||
reason: str,
|
||||
*,
|
||||
fallback: bool,
|
||||
cloud_suppressed_by_local_mode: bool = False,
|
||||
technical_fallback: bool = False,
|
||||
) -> Iterator[str]:
|
||||
if fallback:
|
||||
self._notice_callback(f"{reason} – lokale Antwort mit Ollama.")
|
||||
emitted = False
|
||||
try:
|
||||
for chunk in self._provider_chunks(self.local_provider, messages):
|
||||
emitted = True
|
||||
yield chunk
|
||||
except ProviderError as exc:
|
||||
self._record(
|
||||
ProviderEvent(
|
||||
provider=self.local_provider.name,
|
||||
model=self.local_provider.model,
|
||||
success=False,
|
||||
error_category=type(exc).__name__,
|
||||
fallback=fallback,
|
||||
privacy_policy=decision.policy.value,
|
||||
)
|
||||
)
|
||||
raise
|
||||
if not emitted:
|
||||
raise InvalidProviderResponseError(
|
||||
"Ollama hat keine verwendbare Streaming-Antwort geliefert."
|
||||
)
|
||||
self._record(
|
||||
ProviderEvent(
|
||||
provider=self.local_provider.name,
|
||||
model=self.local_provider.model,
|
||||
success=True,
|
||||
error_category=None,
|
||||
fallback=fallback,
|
||||
privacy_policy=decision.policy.value,
|
||||
)
|
||||
)
|
||||
self.last_route = RouteStatus(
|
||||
provider=self.local_provider.name,
|
||||
privacy_policy=decision.policy,
|
||||
fallback=fallback,
|
||||
cloud_suppressed_by_local_mode=cloud_suppressed_by_local_mode,
|
||||
technical_fallback=technical_fallback,
|
||||
reason=reason,
|
||||
)
|
||||
|
||||
def _cloud_unavailable_reason(self) -> str | None:
|
||||
if not self._cloud_enabled:
|
||||
return "Gemini ist lokal nicht aktiviert"
|
||||
@@ -188,17 +429,7 @@ class HybridProvider:
|
||||
*,
|
||||
disable_notice: bool = False,
|
||||
) -> str:
|
||||
error_category = type(error).__name__
|
||||
self._record(
|
||||
ProviderEvent(
|
||||
provider="gemini",
|
||||
model=self._cloud_model,
|
||||
success=False,
|
||||
error_category=error_category,
|
||||
fallback=False,
|
||||
privacy_policy=decision.policy.value,
|
||||
)
|
||||
)
|
||||
self._record_cloud_error(decision, error)
|
||||
reason = "Gemini nicht verfügbar"
|
||||
if disable_notice:
|
||||
reason = "Gemini-Schlüssel abgelehnt; Cloud für diese Sitzung deaktiviert"
|
||||
@@ -210,6 +441,22 @@ class HybridProvider:
|
||||
technical_fallback=True,
|
||||
)
|
||||
|
||||
def _record_cloud_error(
|
||||
self,
|
||||
decision: PrivacyDecision,
|
||||
error: ProviderError,
|
||||
) -> None:
|
||||
self._record(
|
||||
ProviderEvent(
|
||||
provider="gemini",
|
||||
model=self._cloud_model,
|
||||
success=False,
|
||||
error_category=type(error).__name__,
|
||||
fallback=False,
|
||||
privacy_policy=decision.policy.value,
|
||||
)
|
||||
)
|
||||
|
||||
def _local(
|
||||
self,
|
||||
messages: list[ChatMessage],
|
||||
|
||||
@@ -4,6 +4,7 @@ from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import getpass
|
||||
import re
|
||||
import sys
|
||||
from collections.abc import Callable, Sequence
|
||||
from pathlib import Path
|
||||
@@ -27,6 +28,7 @@ from javis.security.secrets import SecretProvider, SecretStoreError
|
||||
|
||||
InputFunction = Callable[[str], str]
|
||||
OutputFunction = Callable[[str], None]
|
||||
StreamOutputFunction = Callable[[str], None]
|
||||
StatusFunction = Callable[[], list[str]]
|
||||
|
||||
HELP_TEXT = """Befehle:
|
||||
@@ -144,12 +146,48 @@ def _show_sessions(sessions: list[ChatSession], output: OutputFunction) -> None:
|
||||
)
|
||||
|
||||
|
||||
class TerminalMarkdownRenderer:
|
||||
"""Turn common Markdown markers into readable incremental terminal text."""
|
||||
|
||||
_tail_size = 3
|
||||
|
||||
def __init__(self) -> None:
|
||||
self._pending = ""
|
||||
|
||||
@staticmethod
|
||||
def _clean(text: str) -> str:
|
||||
text = re.sub(r"(?m)^[ \t]{0,3}#{1,6}[ \t]+", "", text)
|
||||
for marker in ("```", "**", "__", "~~", "`"):
|
||||
text = text.replace(marker, "")
|
||||
return text
|
||||
|
||||
def feed(self, chunk: str) -> str:
|
||||
combined = self._pending + chunk
|
||||
if len(combined) <= self._tail_size:
|
||||
self._pending = combined
|
||||
return ""
|
||||
visible = combined[: -self._tail_size]
|
||||
self._pending = combined[-self._tail_size :]
|
||||
return self._clean(visible)
|
||||
|
||||
def finish(self) -> str:
|
||||
visible = self._clean(self._pending)
|
||||
self._pending = ""
|
||||
return visible
|
||||
|
||||
|
||||
def _terminal_write(text: str) -> None:
|
||||
sys.stdout.write(text)
|
||||
sys.stdout.flush()
|
||||
|
||||
|
||||
def run_chat(
|
||||
service: ChatService,
|
||||
*,
|
||||
session_id: str | None = None,
|
||||
input_fn: InputFunction = input,
|
||||
output: OutputFunction = print,
|
||||
stream_output: StreamOutputFunction | None = None,
|
||||
status_fn: StatusFunction | None = None,
|
||||
) -> int:
|
||||
try:
|
||||
@@ -262,12 +300,41 @@ def run_chat(
|
||||
output("Unbekannter Befehl. /help zeigt die verfügbaren Befehle.")
|
||||
continue
|
||||
|
||||
stream = service.stream_send(active_id, entered)
|
||||
writer = stream_output or (_terminal_write if output is print else output)
|
||||
renderer = TerminalMarkdownRenderer()
|
||||
started = False
|
||||
try:
|
||||
response = service.send(active_id, entered)
|
||||
output(f"Javis: {response}")
|
||||
for chunk in stream:
|
||||
visible = renderer.feed(chunk)
|
||||
if not visible:
|
||||
continue
|
||||
if not started:
|
||||
writer("Javis: ")
|
||||
started = True
|
||||
writer(visible)
|
||||
remaining = renderer.finish()
|
||||
if remaining:
|
||||
if not started:
|
||||
writer("Javis: ")
|
||||
started = True
|
||||
writer(remaining)
|
||||
if started:
|
||||
writer("\n")
|
||||
except KeyboardInterrupt:
|
||||
close = getattr(stream, "close", None)
|
||||
if callable(close):
|
||||
close()
|
||||
if started:
|
||||
writer("\n")
|
||||
output("Generierung abgebrochen. Die unvollständige Antwort wurde nicht gespeichert.")
|
||||
except (ProviderError, SessionStoreError, SessionProviderMismatchError) as exc:
|
||||
if started:
|
||||
writer("\n")
|
||||
output(f"Fehler: {exc}")
|
||||
except ValueError as exc:
|
||||
if started:
|
||||
writer("\n")
|
||||
output(f"Fehler: {exc}")
|
||||
|
||||
|
||||
|
||||
@@ -2,6 +2,7 @@
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Iterator
|
||||
from dataclasses import dataclass
|
||||
from typing import Protocol, runtime_checkable
|
||||
|
||||
@@ -70,3 +71,13 @@ class LocalModelProvider(Protocol):
|
||||
def chat(self, messages: list[ChatMessage]) -> str:
|
||||
"""Return one complete assistant response."""
|
||||
...
|
||||
|
||||
|
||||
@runtime_checkable
|
||||
class StreamingModelProvider(Protocol):
|
||||
name: str
|
||||
model: str
|
||||
|
||||
def stream_chat(self, messages: list[ChatMessage]) -> Iterator[str]:
|
||||
"""Yield only visible assistant text chunks."""
|
||||
...
|
||||
|
||||
@@ -2,7 +2,7 @@
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Callable
|
||||
from collections.abc import Callable, Iterator
|
||||
from typing import Any
|
||||
|
||||
import httpx
|
||||
@@ -49,6 +49,7 @@ class GeminiProvider:
|
||||
self._config = types.GenerateContentConfig(
|
||||
candidate_count=1,
|
||||
max_output_tokens=max_output_tokens,
|
||||
thinking_config=types.ThinkingConfig(include_thoughts=False),
|
||||
)
|
||||
self._client = client_factory(
|
||||
api_key=api_key,
|
||||
@@ -104,6 +105,73 @@ class GeminiProvider:
|
||||
)
|
||||
return text.strip()
|
||||
|
||||
def stream_chat(self, messages: list[ChatMessage]) -> Iterator[str]:
|
||||
contents = [
|
||||
types.Content(
|
||||
role="model" if message.role == "assistant" else "user",
|
||||
parts=[types.Part.from_text(text=message.content)],
|
||||
)
|
||||
for message in messages
|
||||
]
|
||||
self.last_usage = ProviderUsage()
|
||||
visible_text_received = False
|
||||
try:
|
||||
responses = self._client.models.generate_content_stream(
|
||||
model=self.model,
|
||||
contents=contents,
|
||||
config=self._config,
|
||||
)
|
||||
for response in responses:
|
||||
usage = getattr(response, "usage_metadata", None)
|
||||
if usage is not None:
|
||||
self.last_usage = ProviderUsage(
|
||||
input_tokens=self._token_count(usage, "prompt_token_count"),
|
||||
output_tokens=self._token_count(
|
||||
usage,
|
||||
"candidates_token_count",
|
||||
),
|
||||
)
|
||||
for text in self._visible_text_parts(response):
|
||||
visible_text_received = True
|
||||
yield text
|
||||
except errors.APIError as exc:
|
||||
self._raise_api_error(exc)
|
||||
except httpx.TimeoutException as exc:
|
||||
raise ProviderTimeoutError("Gemini hat nicht rechtzeitig geantwortet.") from exc
|
||||
except httpx.NetworkError as exc:
|
||||
raise CloudNetworkError(
|
||||
"Gemini ist wegen eines Netzwerkfehlers nicht erreichbar."
|
||||
) from exc
|
||||
except (OSError, ConnectionError) as exc:
|
||||
raise CloudNetworkError(
|
||||
"Gemini ist wegen eines Netzwerkfehlers nicht erreichbar."
|
||||
) from exc
|
||||
|
||||
if not visible_text_received:
|
||||
raise InvalidProviderResponseError(
|
||||
"Gemini hat keine verwendbare Streaming-Antwort geliefert."
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _visible_text_parts(response: object) -> Iterator[str]:
|
||||
candidates = getattr(response, "candidates", None)
|
||||
if isinstance(candidates, list) and candidates:
|
||||
for candidate in candidates:
|
||||
content = getattr(candidate, "content", None)
|
||||
parts = getattr(content, "parts", None)
|
||||
if not isinstance(parts, list):
|
||||
continue
|
||||
for part in parts:
|
||||
if getattr(part, "thought", False):
|
||||
continue
|
||||
text = getattr(part, "text", None)
|
||||
if isinstance(text, str) and text:
|
||||
yield text
|
||||
return
|
||||
text = getattr(response, "text", None)
|
||||
if isinstance(text, str) and text:
|
||||
yield text
|
||||
|
||||
@staticmethod
|
||||
def _token_count(usage: object, name: str) -> int | None:
|
||||
value = getattr(usage, name, None)
|
||||
|
||||
@@ -4,6 +4,7 @@ from __future__ import annotations
|
||||
|
||||
import json
|
||||
import socket
|
||||
from collections.abc import Iterator
|
||||
from dataclasses import dataclass
|
||||
from urllib.error import HTTPError, URLError
|
||||
from urllib.request import Request, urlopen
|
||||
@@ -107,3 +108,76 @@ class OllamaProvider:
|
||||
"Ollama hat keine verwendbare Textantwort geliefert."
|
||||
)
|
||||
return content.strip()
|
||||
|
||||
def stream_chat(self, messages: list[ChatMessage]) -> Iterator[str]:
|
||||
payload = {
|
||||
"model": self.model,
|
||||
"messages": [
|
||||
{"role": message.role, "content": message.content} for message in messages
|
||||
],
|
||||
"stream": True,
|
||||
"think": False,
|
||||
}
|
||||
request = Request(
|
||||
self._endpoint,
|
||||
data=json.dumps(payload).encode("utf-8"),
|
||||
headers={"Content-Type": "application/json"},
|
||||
method="POST",
|
||||
)
|
||||
completed = False
|
||||
visible_text_received = False
|
||||
try:
|
||||
with urlopen(request, timeout=self._timeout_seconds) as response:
|
||||
for raw_line in response:
|
||||
if not raw_line.strip():
|
||||
continue
|
||||
try:
|
||||
result = json.loads(raw_line)
|
||||
except (json.JSONDecodeError, UnicodeDecodeError) as exc:
|
||||
raise InvalidProviderResponseError(
|
||||
"Ollama hat einen ungültigen Streaming-Abschnitt geliefert."
|
||||
) from exc
|
||||
error = result.get("error")
|
||||
if isinstance(error, str) and error:
|
||||
if "not found" in error.lower():
|
||||
raise ModelNotInstalledError(
|
||||
f"Das lokale Modell '{self.model}' ist nicht installiert."
|
||||
)
|
||||
raise ProviderUnavailableError("Ollama hat den Stream abgelehnt.")
|
||||
content = result.get("message", {}).get("content")
|
||||
if content is not None and not isinstance(content, str):
|
||||
raise InvalidProviderResponseError(
|
||||
"Ollama hat ungültigen sichtbaren Text geliefert."
|
||||
)
|
||||
if content:
|
||||
visible_text_received = True
|
||||
yield content
|
||||
if result.get("done") is True:
|
||||
completed = True
|
||||
break
|
||||
except HTTPError as exc:
|
||||
details = exc.read().decode("utf-8", errors="replace")
|
||||
if exc.code == 404 or "not found" in details.lower():
|
||||
raise ModelNotInstalledError(
|
||||
f"Das lokale Modell '{self.model}' ist nicht installiert."
|
||||
) from exc
|
||||
raise ProviderUnavailableError(f"Ollama meldet HTTP-Fehler {exc.code}.") from exc
|
||||
except TimeoutError as exc:
|
||||
raise ProviderTimeoutError(
|
||||
"Die Modellantwort hat das Zeitlimit überschritten."
|
||||
) from exc
|
||||
except URLError as exc:
|
||||
if isinstance(exc.reason, (TimeoutError, socket.timeout)):
|
||||
raise ProviderTimeoutError(
|
||||
"Die Modellantwort hat das Zeitlimit überschritten."
|
||||
) from exc
|
||||
raise ProviderUnavailableError(
|
||||
"Ollama ist unter der konfigurierten lokalen Adresse nicht erreichbar."
|
||||
) from exc
|
||||
|
||||
if not completed:
|
||||
raise ResponseAbortedError("Der Ollama-Stream wurde vorzeitig beendet.")
|
||||
if not visible_text_received:
|
||||
raise InvalidProviderResponseError(
|
||||
"Ollama hat keine verwendbare Streaming-Antwort geliefert."
|
||||
)
|
||||
|
||||
@@ -4,7 +4,7 @@ from pathlib import Path
|
||||
|
||||
from javis.core.chat_service import ChatService
|
||||
from javis.memory.sqlite_store import SQLiteSessionStore
|
||||
from javis.providers.base import ChatMessage, LocalModelProvider
|
||||
from javis.providers.base import ChatMessage, LocalModelProvider, ProviderUnavailableError
|
||||
|
||||
|
||||
class RecordingProvider:
|
||||
@@ -19,6 +19,19 @@ class RecordingProvider:
|
||||
return f"Antwort {len(self.calls)}"
|
||||
|
||||
|
||||
class StreamingProvider(RecordingProvider):
|
||||
def __init__(self, *, fail_after_first: bool = False) -> None:
|
||||
super().__init__()
|
||||
self.fail_after_first = fail_after_first
|
||||
|
||||
def stream_chat(self, messages: list[ChatMessage]):
|
||||
self.calls.append(messages)
|
||||
yield "Teil "
|
||||
if self.fail_after_first:
|
||||
raise ProviderUnavailableError("Streamfehler")
|
||||
yield "Antwort"
|
||||
|
||||
|
||||
class ChatServiceTests(unittest.TestCase):
|
||||
def setUp(self) -> None:
|
||||
self.temporary_directory = tempfile.TemporaryDirectory()
|
||||
@@ -52,6 +65,41 @@ class ChatServiceTests(unittest.TestCase):
|
||||
],
|
||||
)
|
||||
|
||||
def test_streaming_response_is_persisted_exactly_once_after_completion(self) -> None:
|
||||
provider = StreamingProvider()
|
||||
service = ChatService(self.service.store, provider)
|
||||
session = service.new_session()
|
||||
|
||||
chunks = list(service.stream_send(session.id, "Hallo"))
|
||||
|
||||
self.assertEqual(chunks, ["Teil ", "Antwort"])
|
||||
self.assertEqual(
|
||||
[message.content for message in service.load_session(session.id).messages],
|
||||
["Hallo", "Teil Antwort"],
|
||||
)
|
||||
self.assertEqual(service.load_session(session.id).session.message_count, 2)
|
||||
|
||||
def test_closed_stream_does_not_persist_partial_response(self) -> None:
|
||||
provider = StreamingProvider()
|
||||
service = ChatService(self.service.store, provider)
|
||||
session = service.new_session()
|
||||
stream = service.stream_send(session.id, "Hallo")
|
||||
|
||||
self.assertEqual(next(stream), "Teil ")
|
||||
stream.close()
|
||||
|
||||
self.assertEqual(service.load_session(session.id).messages, [])
|
||||
|
||||
def test_provider_error_during_stream_does_not_persist_partial_response(self) -> None:
|
||||
provider = StreamingProvider(fail_after_first=True)
|
||||
service = ChatService(self.service.store, provider)
|
||||
session = service.new_session()
|
||||
|
||||
with self.assertRaises(ProviderUnavailableError):
|
||||
list(service.stream_send(session.id, "Hallo"))
|
||||
|
||||
self.assertEqual(service.load_session(session.id).messages, [])
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
|
||||
+60
-1
@@ -19,6 +19,22 @@ class _FakeProvider:
|
||||
def chat(self, messages: list[ChatMessage]) -> str:
|
||||
return f"Echo: {messages[-1].content}"
|
||||
|
||||
def stream_chat(self, messages: list[ChatMessage]):
|
||||
yield "Echo: "
|
||||
yield messages[-1].content
|
||||
|
||||
|
||||
class _MarkdownProvider(_FakeProvider):
|
||||
def stream_chat(self, messages: list[ChatMessage]):
|
||||
yield "## **Ant"
|
||||
yield "wort**\n`SQLite`"
|
||||
|
||||
|
||||
class _InterruptingProvider(_FakeProvider):
|
||||
def stream_chat(self, messages: list[ChatMessage]):
|
||||
yield "angefangene Antwort"
|
||||
raise KeyboardInterrupt
|
||||
|
||||
|
||||
class _FakeHybridProvider(_FakeProvider):
|
||||
name = "hybrid"
|
||||
@@ -59,10 +75,53 @@ class CliTests(unittest.TestCase):
|
||||
)
|
||||
|
||||
self.assertEqual(result, 0)
|
||||
self.assertTrue(any(line == "Javis: Echo: Hallo" for line in output))
|
||||
self.assertIn("Javis: Echo: Hallo\n", "".join(output))
|
||||
self.assertTrue(any("Nachrichten" in line for line in output))
|
||||
self.assertEqual(output[-1], "Chat beendet.")
|
||||
|
||||
def test_streaming_cleans_terminal_markdown_but_persists_raw_answer(self) -> None:
|
||||
with tempfile.TemporaryDirectory() as directory:
|
||||
store = SQLiteSessionStore(Path(directory) / "sessions.sqlite3")
|
||||
service = ChatService(store, _MarkdownProvider())
|
||||
inputs = iter(["Frage", "/exit"])
|
||||
output: list[str] = []
|
||||
streamed: list[str] = []
|
||||
|
||||
result = run_chat(
|
||||
service,
|
||||
input_fn=lambda _prompt: next(inputs),
|
||||
output=output.append,
|
||||
stream_output=streamed.append,
|
||||
)
|
||||
session = service.list_sessions()[0]
|
||||
messages = service.load_session(session.id).messages
|
||||
|
||||
self.assertEqual(result, 0)
|
||||
self.assertEqual("".join(streamed), "Javis: Antwort\nSQLite\n")
|
||||
self.assertEqual(messages[-1].content, "## **Antwort**\n`SQLite`")
|
||||
|
||||
def test_keyboard_interrupt_drops_partial_answer_and_returns_to_prompt(self) -> None:
|
||||
with tempfile.TemporaryDirectory() as directory:
|
||||
store = SQLiteSessionStore(Path(directory) / "sessions.sqlite3")
|
||||
service = ChatService(store, _InterruptingProvider())
|
||||
inputs = iter(["Frage", "/exit"])
|
||||
output: list[str] = []
|
||||
streamed: list[str] = []
|
||||
|
||||
result = run_chat(
|
||||
service,
|
||||
input_fn=lambda _prompt: next(inputs),
|
||||
output=output.append,
|
||||
stream_output=streamed.append,
|
||||
)
|
||||
session = service.list_sessions()[0]
|
||||
messages = service.load_session(session.id).messages
|
||||
|
||||
self.assertEqual(result, 0)
|
||||
self.assertEqual(messages, [])
|
||||
self.assertTrue(any("nicht gespeichert" in line for line in output))
|
||||
self.assertEqual(output[-1], "Chat beendet.")
|
||||
|
||||
def test_provider_privacy_and_status_commands_are_sanitized(self) -> None:
|
||||
with tempfile.TemporaryDirectory() as directory:
|
||||
store = SQLiteSessionStore(Path(directory) / "sessions.sqlite3")
|
||||
|
||||
@@ -23,6 +23,7 @@ class _FakeModels:
|
||||
def __init__(self, result: object) -> None:
|
||||
self.result = result
|
||||
self.call: dict[str, Any] | None = None
|
||||
self.stream_call: dict[str, Any] | None = None
|
||||
|
||||
def generate_content(self, **kwargs: object) -> object:
|
||||
self.call = kwargs
|
||||
@@ -30,6 +31,14 @@ class _FakeModels:
|
||||
raise self.result
|
||||
return self.result
|
||||
|
||||
def generate_content_stream(self, **kwargs: object):
|
||||
self.stream_call = kwargs
|
||||
if isinstance(self.result, BaseException):
|
||||
raise self.result
|
||||
if isinstance(self.result, list):
|
||||
return iter(self.result)
|
||||
return iter((self.result,))
|
||||
|
||||
|
||||
class _FakeClient:
|
||||
def __init__(self, result: object) -> None:
|
||||
@@ -56,6 +65,23 @@ def _response(text: str = "Cloud-Antwort") -> SimpleNamespace:
|
||||
)
|
||||
|
||||
|
||||
def _stream_response(
|
||||
visible_text: str,
|
||||
*,
|
||||
thought_text: str | None = None,
|
||||
with_usage: bool = False,
|
||||
) -> SimpleNamespace:
|
||||
parts = []
|
||||
if thought_text:
|
||||
parts.append(SimpleNamespace(text=thought_text, thought=True))
|
||||
parts.append(SimpleNamespace(text=visible_text, thought=False))
|
||||
usage = SimpleNamespace(prompt_token_count=12, candidates_token_count=7) if with_usage else None
|
||||
return SimpleNamespace(
|
||||
candidates=[SimpleNamespace(content=SimpleNamespace(parts=parts))],
|
||||
usage_metadata=usage,
|
||||
)
|
||||
|
||||
|
||||
class GeminiProviderTests(unittest.TestCase):
|
||||
def _provider(
|
||||
self, result: object, **overrides: object
|
||||
@@ -99,6 +125,39 @@ class GeminiProviderTests(unittest.TestCase):
|
||||
self.assertEqual([content.role for content in call["contents"]], ["user", "model", "user"])
|
||||
self.assertEqual(call["config"].max_output_tokens, 1024)
|
||||
self.assertIsNone(call["config"].tools)
|
||||
self.assertFalse(call["config"].thinking_config.include_thoughts)
|
||||
|
||||
def test_streaming_yields_only_visible_text_and_tracks_usage(self) -> None:
|
||||
provider, factory = self._provider(
|
||||
[
|
||||
_stream_response("Cloud ", thought_text="verstecktes Denken"),
|
||||
_stream_response("Antwort", with_usage=True),
|
||||
]
|
||||
)
|
||||
|
||||
chunks = list(provider.stream_chat([ChatMessage("user", "Was ist SQLite?")]))
|
||||
|
||||
self.assertEqual(chunks, ["Cloud ", "Antwort"])
|
||||
self.assertNotIn("verstecktes Denken", "".join(chunks))
|
||||
self.assertEqual(provider.last_usage.input_tokens, 12)
|
||||
self.assertEqual(provider.last_usage.output_tokens, 7)
|
||||
self.assertEqual(
|
||||
factory.client.models.stream_call["model"],
|
||||
"gemini-3.6-flash",
|
||||
)
|
||||
|
||||
def test_streaming_error_is_sanitized(self) -> None:
|
||||
provider, _ = self._provider(
|
||||
errors.ClientError(
|
||||
429,
|
||||
{"message": "quota test-key-not-a-real-secret"},
|
||||
)
|
||||
)
|
||||
|
||||
with self.assertRaises(ProviderRateLimitError) as raised:
|
||||
list(provider.stream_chat([ChatMessage("user", "Was ist Python?")]))
|
||||
|
||||
self.assertNotIn("test-key-not-a-real-secret", str(raised.exception))
|
||||
|
||||
def test_missing_key_stops_before_client_creation(self) -> None:
|
||||
factory = _ClientFactory(_response())
|
||||
|
||||
@@ -4,22 +4,39 @@ import unittest
|
||||
from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer
|
||||
from typing import ClassVar
|
||||
|
||||
from javis.providers.base import ChatMessage
|
||||
from javis.providers.base import ChatMessage, ResponseAbortedError
|
||||
from javis.providers.ollama import OllamaProvider
|
||||
|
||||
|
||||
class _OllamaHandler(BaseHTTPRequestHandler):
|
||||
request_payload: ClassVar[dict[str, object]] = {}
|
||||
abort_stream: ClassVar[bool] = False
|
||||
|
||||
def do_POST(self) -> None:
|
||||
length = int(self.headers["Content-Length"])
|
||||
type(self).request_payload = json.loads(self.rfile.read(length))
|
||||
body = json.dumps(
|
||||
{
|
||||
"message": {"role": "assistant", "content": "Lokale Antwort"},
|
||||
"done": True,
|
||||
}
|
||||
).encode()
|
||||
if type(self).request_payload["stream"]:
|
||||
chunks = [
|
||||
{
|
||||
"message": {"role": "assistant", "content": "Lokale "},
|
||||
"done": False,
|
||||
}
|
||||
]
|
||||
if not type(self).abort_stream:
|
||||
chunks.append(
|
||||
{
|
||||
"message": {"role": "assistant", "content": "Antwort"},
|
||||
"done": True,
|
||||
}
|
||||
)
|
||||
body = b"".join(json.dumps(chunk).encode() + b"\n" for chunk in chunks)
|
||||
else:
|
||||
body = json.dumps(
|
||||
{
|
||||
"message": {"role": "assistant", "content": "Lokale Antwort"},
|
||||
"done": True,
|
||||
}
|
||||
).encode()
|
||||
self.send_response(200)
|
||||
self.send_header("Content-Type", "application/json")
|
||||
self.send_header("Content-Length", str(len(body)))
|
||||
@@ -39,6 +56,9 @@ class _OllamaHandler(BaseHTTPRequestHandler):
|
||||
|
||||
|
||||
class OllamaProviderTests(unittest.TestCase):
|
||||
def setUp(self) -> None:
|
||||
_OllamaHandler.abort_stream = False
|
||||
|
||||
def test_provider_uses_local_chat_endpoint(self) -> None:
|
||||
server = ThreadingHTTPServer(("127.0.0.1", 0), _OllamaHandler)
|
||||
thread = threading.Thread(target=server.serve_forever, daemon=True)
|
||||
@@ -62,6 +82,44 @@ class OllamaProviderTests(unittest.TestCase):
|
||||
self.assertTrue(status.reachable)
|
||||
self.assertTrue(status.model_available)
|
||||
|
||||
def test_provider_streams_visible_text_without_thinking(self) -> None:
|
||||
server = ThreadingHTTPServer(("127.0.0.1", 0), _OllamaHandler)
|
||||
thread = threading.Thread(target=server.serve_forever, daemon=True)
|
||||
thread.start()
|
||||
try:
|
||||
provider = OllamaProvider(
|
||||
"test-model",
|
||||
f"http://127.0.0.1:{server.server_port}",
|
||||
2,
|
||||
)
|
||||
chunks = list(provider.stream_chat([ChatMessage("user", "Hallo")]))
|
||||
finally:
|
||||
server.shutdown()
|
||||
server.server_close()
|
||||
thread.join()
|
||||
|
||||
self.assertEqual(chunks, ["Lokale ", "Antwort"])
|
||||
self.assertTrue(_OllamaHandler.request_payload["stream"])
|
||||
self.assertFalse(_OllamaHandler.request_payload["think"])
|
||||
|
||||
def test_provider_rejects_incomplete_stream(self) -> None:
|
||||
_OllamaHandler.abort_stream = True
|
||||
server = ThreadingHTTPServer(("127.0.0.1", 0), _OllamaHandler)
|
||||
thread = threading.Thread(target=server.serve_forever, daemon=True)
|
||||
thread.start()
|
||||
try:
|
||||
provider = OllamaProvider(
|
||||
"test-model",
|
||||
f"http://127.0.0.1:{server.server_port}",
|
||||
2,
|
||||
)
|
||||
with self.assertRaises(ResponseAbortedError):
|
||||
list(provider.stream_chat([ChatMessage("user", "Hallo")]))
|
||||
finally:
|
||||
server.shutdown()
|
||||
server.server_close()
|
||||
thread.join()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
|
||||
@@ -16,6 +16,7 @@ from javis.providers.base import (
|
||||
ProviderRateLimitError,
|
||||
ProviderTimeoutError,
|
||||
ProviderUsage,
|
||||
ResponseAbortedError,
|
||||
)
|
||||
from javis.security.privacy import CloudPolicy, PrivacyRouter
|
||||
|
||||
@@ -28,11 +29,15 @@ class _RecordingProvider:
|
||||
*,
|
||||
answer: str = "Antwort",
|
||||
error: Exception | None = None,
|
||||
stream_chunks: tuple[str, ...] | None = None,
|
||||
stream_error: Exception | None = None,
|
||||
) -> None:
|
||||
self.name = name
|
||||
self.model = model
|
||||
self.answer = answer
|
||||
self.error = error
|
||||
self.stream_chunks = stream_chunks
|
||||
self.stream_error = stream_error
|
||||
self.calls: list[list[ChatMessage]] = []
|
||||
self.last_usage = ProviderUsage(11, 5)
|
||||
|
||||
@@ -42,6 +47,17 @@ class _RecordingProvider:
|
||||
raise self.error
|
||||
return self.answer
|
||||
|
||||
def stream_chat(self, messages: list[ChatMessage]):
|
||||
self.calls.append(messages)
|
||||
if self.stream_chunks is None:
|
||||
if self.error:
|
||||
raise self.error
|
||||
yield self.answer
|
||||
else:
|
||||
yield from self.stream_chunks
|
||||
if self.stream_error:
|
||||
raise self.stream_error
|
||||
|
||||
|
||||
class HybridProviderTests(unittest.TestCase):
|
||||
def setUp(self) -> None:
|
||||
@@ -283,6 +299,64 @@ class HybridProviderTests(unittest.TestCase):
|
||||
self.assertFalse(router.last_route.cloud_suppressed_by_local_mode)
|
||||
self.assertTrue(router.last_route.technical_fallback)
|
||||
|
||||
def test_allowed_cloud_response_streams_visible_chunks(self) -> None:
|
||||
self.cloud.stream_chunks = ("Cloud ", "Stream")
|
||||
router = self._router()
|
||||
|
||||
chunks = list(router.stream_chat([ChatMessage("user", "Wie funktioniert SQLite?")]))
|
||||
|
||||
self.assertEqual(chunks, ["Cloud ", "Stream"])
|
||||
self.assertEqual(router.last_route.provider, "gemini")
|
||||
self.assertFalse(router.last_route.fallback)
|
||||
|
||||
def test_cloud_stream_error_before_output_falls_back_locally(self) -> None:
|
||||
self.cloud.stream_chunks = ()
|
||||
self.cloud.stream_error = CloudNetworkError("offline")
|
||||
self.local.stream_chunks = ("Lokaler ", "Ersatz")
|
||||
router = self._router()
|
||||
|
||||
chunks = list(router.stream_chat([ChatMessage("user", "Wie funktioniert Python?")]))
|
||||
|
||||
self.assertEqual(chunks, ["Lokaler ", "Ersatz"])
|
||||
self.assertTrue(router.last_route.fallback)
|
||||
self.assertTrue(router.last_route.technical_fallback)
|
||||
self.assertIn("lokale Antwort", self.notices[-1])
|
||||
|
||||
def test_cloud_stream_error_after_output_never_appends_local_answer(self) -> None:
|
||||
self.cloud.stream_chunks = ("Teilantwort",)
|
||||
self.cloud.stream_error = CloudNetworkError("offline")
|
||||
self.local.stream_chunks = ("Lokaler Ersatz",)
|
||||
router = self._router()
|
||||
stream = router.stream_chat([ChatMessage("user", "Wie funktioniert Python?")])
|
||||
|
||||
self.assertEqual(next(stream), "Teilantwort")
|
||||
with self.assertRaises(ResponseAbortedError):
|
||||
next(stream)
|
||||
|
||||
self.assertFalse(self.local.calls)
|
||||
self.assertFalse(router.last_route.fallback)
|
||||
self.assertTrue(router.last_route.technical_fallback)
|
||||
|
||||
def test_medical_never_streams_locally_before_cloud_construction(self) -> None:
|
||||
constructed = False
|
||||
self.local.stream_chunks = ("Lokal",)
|
||||
|
||||
def cloud_factory() -> _RecordingProvider:
|
||||
nonlocal constructed
|
||||
constructed = True
|
||||
return self.cloud
|
||||
|
||||
router = self._router(cloud_provider_factory=cloud_factory, mode="gemini")
|
||||
|
||||
chunks = list(
|
||||
router.stream_chat([ChatMessage("user", "Welche Diagnose passt zu meinen Schmerzen?")])
|
||||
)
|
||||
|
||||
self.assertEqual(chunks, ["Lokal"])
|
||||
self.assertFalse(constructed)
|
||||
self.assertFalse(self.approvals)
|
||||
self.assertEqual(router.last_route.privacy_policy, CloudPolicy.NEVER)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
|
||||
Reference in New Issue
Block a user