From 00a5a6ce8cb6e8783c77bae2d2acc8ee3807b831 Mon Sep 17 00:00:00 2001 From: Dystroyer8 Date: Thu, 30 Jul 2026 19:38:53 +0200 Subject: [PATCH] feat: stream provider responses --- docs/CHATGPT_HANDOFF.md | 33 ++-- src/javis/core/chat_service.py | 45 ++++- src/javis/core/provider_router.py | 271 +++++++++++++++++++++++++++-- src/javis/interface/cli.py | 71 +++++++- src/javis/providers/base.py | 11 ++ src/javis/providers/gemini.py | 70 +++++++- src/javis/providers/ollama.py | 74 ++++++++ tests/unit/test_chat_service.py | 50 +++++- tests/unit/test_cli.py | 61 ++++++- tests/unit/test_gemini_provider.py | 59 +++++++ tests/unit/test_ollama_provider.py | 72 +++++++- tests/unit/test_provider_router.py | 74 ++++++++ 12 files changed, 850 insertions(+), 41 deletions(-) diff --git a/docs/CHATGPT_HANDOFF.md b/docs/CHATGPT_HANDOFF.md index 3f49679..1d0b0d8 100644 --- a/docs/CHATGPT_HANDOFF.md +++ b/docs/CHATGPT_HANDOFF.md @@ -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. diff --git a/src/javis/core/chat_service.py b/src/javis/core/chat_service.py index 5e790e2..4682dec 100644 --- a/src/javis/core/chat_service.py +++ b/src/javis/core/chat_service.py @@ -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) diff --git a/src/javis/core/provider_router.py b/src/javis/core/provider_router.py index 482e4a3..8fddff3 100644 --- a/src/javis/core/provider_router.py +++ b/src/javis/core/provider_router.py @@ -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], diff --git a/src/javis/interface/cli.py b/src/javis/interface/cli.py index 8c618a4..1542a7d 100644 --- a/src/javis/interface/cli.py +++ b/src/javis/interface/cli.py @@ -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}") diff --git a/src/javis/providers/base.py b/src/javis/providers/base.py index 8928ead..bc71a24 100644 --- a/src/javis/providers/base.py +++ b/src/javis/providers/base.py @@ -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.""" + ... diff --git a/src/javis/providers/gemini.py b/src/javis/providers/gemini.py index 7e7ae39..baca8b9 100644 --- a/src/javis/providers/gemini.py +++ b/src/javis/providers/gemini.py @@ -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) diff --git a/src/javis/providers/ollama.py b/src/javis/providers/ollama.py index 474ae86..5c99f7a 100644 --- a/src/javis/providers/ollama.py +++ b/src/javis/providers/ollama.py @@ -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." + ) diff --git a/tests/unit/test_chat_service.py b/tests/unit/test_chat_service.py index 015ed1a..1e3c570 100644 --- a/tests/unit/test_chat_service.py +++ b/tests/unit/test_chat_service.py @@ -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() diff --git a/tests/unit/test_cli.py b/tests/unit/test_cli.py index a8e394c..07ed36d 100644 --- a/tests/unit/test_cli.py +++ b/tests/unit/test_cli.py @@ -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") diff --git a/tests/unit/test_gemini_provider.py b/tests/unit/test_gemini_provider.py index 091a76d..dbcf2d8 100644 --- a/tests/unit/test_gemini_provider.py +++ b/tests/unit/test_gemini_provider.py @@ -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()) diff --git a/tests/unit/test_ollama_provider.py b/tests/unit/test_ollama_provider.py index 46b6d09..388b2b4 100644 --- a/tests/unit/test_ollama_provider.py +++ b/tests/unit/test_ollama_provider.py @@ -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() diff --git a/tests/unit/test_provider_router.py b/tests/unit/test_provider_router.py index c075fd4..4b129eb 100644 --- a/tests/unit/test_provider_router.py +++ b/tests/unit/test_provider_router.py @@ -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()