feat: stream provider responses
This commit is contained in:
@@ -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],
|
||||
|
||||
Reference in New Issue
Block a user