Kansho/backend/providers.py
2026-08-29 20:28:59 +02:00

301 lines
9.8 KiB
Python

"""Two independently configured OpenAI-compatible providers: generate vs detect."""
from __future__ import annotations
import json
import os
import re
from dataclasses import dataclass, field
from urllib.parse import urlparse
import httpx
from env_loader import load_env_file
load_env_file()
CONTEXT_LENGTH_REJECT = re.compile(
r"context.?length|maximum context|prompt is too long|too many tokens|"
r"context window|max context|exceeds? (?:the )?(?:maximum|context)",
re.I,
)
class ProviderError(Exception):
def __init__(self, code: str, message: str, status_code: int = 503, diagnostics: dict | None = None):
super().__init__(message)
self.code = code
self.message = message
self.status_code = status_code
self.diagnostics = diagnostics or {}
PROVIDER_MESSAGE_MAX = 240
_PROVIDER_USAGE_KEYS = ("prompt_tokens", "completion_tokens", "total_tokens", "cost", "total_cost")
def provider_error_diagnostics(status_code: int, body: str, *, extra: dict | None = None) -> dict:
"""Compact provider failure fields. No prompt body, no mapping labels."""
diagnostics: dict = {"http_status": status_code}
if extra:
diagnostics.update({key: value for key, value in extra.items() if value is not None})
parsed = None
try:
parsed = json.loads(body or "")
except (TypeError, ValueError, json.JSONDecodeError):
parsed = None
if isinstance(parsed, dict):
err = parsed.get("error")
if isinstance(err, str) and err.strip():
diagnostics["provider_message"] = err.strip()[:PROVIDER_MESSAGE_MAX]
elif isinstance(err, dict):
msg = str(err.get("message") or err.get("msg") or "").strip()
if msg:
diagnostics["provider_message"] = msg[:PROVIDER_MESSAGE_MAX]
code = err.get("code") or err.get("type")
if code:
diagnostics["provider_code"] = str(code)[:80]
usage = parsed.get("usage")
if isinstance(usage, dict):
kept = {key: usage[key] for key in _PROVIDER_USAGE_KEYS if usage.get(key) is not None}
if kept:
diagnostics["usage"] = kept
if kept.get("cost") is not None:
diagnostics["cost"] = kept["cost"]
elif kept.get("total_cost") is not None:
diagnostics["cost"] = kept["total_cost"]
return diagnostics
def _reject_message(role: str, status_code: int, diagnostics: dict) -> str:
base = f"{role}-Provider hat die Anfrage abgelehnt"
if status_code:
base += f" (HTTP {status_code})"
snippet = (diagnostics.get("provider_message") or "").strip()
if snippet:
return f"{base}: {snippet}"
return f"{base}."
@dataclass(frozen=True)
class ProviderConfig:
role: str
name: str
mode: str
url: str
model: str
key: str
local: bool
zdr: bool
no_train: bool
@dataclass
class ChatResult:
content: str
model: str | None = None
usage: dict = field(default_factory=dict)
context_compression: str = "not_applicable"
finish_reason: str | None = None
def is_openrouter(config: ProviderConfig) -> bool:
host = (urlparse(config.url or "").hostname or "").lower()
return host.endswith("openrouter.ai") or (config.name or "").lower() == "openrouter"
def _truthy(name: str) -> bool:
return os.environ.get(name, "").strip().lower() in {"1", "true", "yes", "on"}
def _is_local_url(url: str) -> bool:
host = (urlparse(url).hostname or "").lower()
return host in {"localhost", "127.0.0.1", "::1", "ollama"}
def _setting(role: str) -> dict:
try:
from provider_settings import get_setting
return get_setting(role) or {}
except Exception:
return {}
def _text(env_name: str, db_value: str | None, default: str = "") -> str:
env = (os.environ.get(env_name) or "").strip()
if env:
return env
if db_value:
return str(db_value).strip()
return default
def _flag(env_name: str, db_value, default: bool = True) -> bool:
if (os.environ.get(env_name) or "").strip():
return _truthy(env_name)
if db_value is None:
return default
return bool(int(db_value) if not isinstance(db_value, bool) else db_value)
def generate_provider() -> ProviderConfig | None:
if _truthy("KANSHO_FAKE_PROVIDER"):
return ProviderConfig(
role="generate",
name="fake",
mode="fake",
url="",
model="fake",
key="",
local=True,
zdr=True,
no_train=True,
)
row = _setting("generate")
key = (os.environ.get("KANSHO_PROVIDER_KEY") or "").strip()
url = _text("KANSHO_PROVIDER_URL", row.get("url"), "https://openrouter.ai/api/v1/chat/completions")
model = _text("KANSHO_PROVIDER_MODEL", row.get("model"), "openai/gpt-4o-mini")
local = _is_local_url(url)
zdr = local or _flag("KANSHO_PROVIDER_ZDR", row.get("zdr"), True)
no_train = local or _flag("KANSHO_PROVIDER_NO_TRAIN", row.get("no_train"), True)
if not local and not key:
return None
if not local and (not zdr or not no_train):
return None
return ProviderConfig(
role="generate",
name=row.get("name") or ("ollama" if local else "openrouter"),
mode="http",
url=url,
model=model,
key=key,
local=local,
zdr=zdr,
no_train=no_train,
)
def detect_provider() -> ProviderConfig | None:
if _truthy("KANSHO_FAKE_DETECT"):
return ProviderConfig(
role="detect",
name="fake-detect",
mode="fake",
url="",
model="fake",
key="",
local=True,
zdr=True,
no_train=True,
)
row = _setting("detect")
url = _text("KANSHO_DETECT_PROVIDER_URL", row.get("url"), "")
if not url:
return None
model = _text("KANSHO_DETECT_PROVIDER_MODEL", row.get("model"), "openai/gpt-4.1-nano")
key = (os.environ.get("KANSHO_DETECT_PROVIDER_KEY") or "").strip()
local = _is_local_url(url)
if not key and not local:
key = (os.environ.get("KANSHO_PROVIDER_KEY") or "").strip()
zdr = local or _flag("KANSHO_DETECT_ZDR", row.get("zdr"), True)
no_train = local or _flag("KANSHO_DETECT_NO_TRAIN", row.get("no_train"), True)
if not local and not key:
return None
if not local and (not zdr or not no_train):
return None
return ProviderConfig(
role="detect",
name=row.get("name") or ("ollama" if local else "openrouter"),
mode="http",
url=url,
model=model,
key=key,
local=local,
zdr=zdr,
no_train=no_train,
)
def complete_chat(
config: ProviderConfig,
messages: list[dict],
*,
timeout: float,
max_tokens: int | None = None,
disable_context_compression: bool = False,
) -> ChatResult:
payload: dict = {"model": config.model, "messages": messages}
if max_tokens is not None:
payload["max_tokens"] = max_tokens
compression = "not_applicable"
if not config.local:
payload["provider"] = {"data_collection": "deny"}
if disable_context_compression and is_openrouter(config) and not config.local:
payload["plugins"] = [{"id": "context-compression", "enabled": False}]
compression = "disabled"
headers = {"Content-Type": "application/json"}
if config.key:
headers["Authorization"] = f"Bearer {config.key}"
if not config.local:
headers["HTTP-Referer"] = "https://kansho.local"
headers["X-Title"] = "Kansho"
try:
response = httpx.post(config.url, json=payload, headers=headers, timeout=timeout)
except httpx.TimeoutException as exc:
raise ProviderError(
"provider_timeout",
f"{config.role}-Provider hat nicht innerhalb von {int(timeout)} Sekunden geantwortet.",
504,
diagnostics={"timeout_s": timeout, "generate_requested": True},
) from exc
except httpx.HTTPError as exc:
raise ProviderError(
"provider_unreachable",
f"{config.role}-Provider ist nicht erreichbar.",
503,
diagnostics={"generate_requested": True},
) from exc
if response.status_code >= 400:
body = ""
try:
body = response.text or ""
except Exception:
body = ""
diagnostics = provider_error_diagnostics(
response.status_code,
body,
extra={"context_compression": compression, "generate_requested": True},
)
if CONTEXT_LENGTH_REJECT.search(body):
raise ProviderError(
"provider_context_length_rejected",
"Der Anbieter hat die Anfrage wegen der Kontextlänge abgelehnt. "
"Kanshō hat den Tagesdialog nicht automatisch gekürzt.",
502,
diagnostics=diagnostics,
)
raise ProviderError(
"provider_rejected",
_reject_message(config.role, response.status_code, diagnostics),
502,
diagnostics=diagnostics,
)
data = response.json()
try:
content = data["choices"][0]["message"]["content"]
except (KeyError, IndexError, TypeError) as exc:
raise ProviderError("provider_shape", f"Unerwartete {config.role}-Antwort.") from exc
usage = data.get("usage") if isinstance(data.get("usage"), dict) else {}
finish_reason = None
try:
finish_reason = data["choices"][0].get("finish_reason")
except (KeyError, IndexError, TypeError, AttributeError):
finish_reason = None
return ChatResult(
content=content or "",
model=data.get("model") or config.model,
usage=usage,
context_compression=compression,
finish_reason=finish_reason,
)