129 lines
6.5 KiB
Python
129 lines
6.5 KiB
Python
import hashlib
|
|
import json
|
|
import math
|
|
from urllib.parse import urlsplit
|
|
|
|
import httpx
|
|
|
|
|
|
class ProviderError(Exception):
|
|
pass
|
|
|
|
|
|
def validate_url(value):
|
|
if not value:
|
|
return ''
|
|
parsed = urlsplit(value)
|
|
if parsed.scheme not in ('http', 'https') or not parsed.hostname or parsed.username or parsed.password or parsed.query or parsed.fragment:
|
|
raise ValueError('Endpoint muss eine HTTP(S)-Basis-URL ohne Zugangsdaten oder Query sein.')
|
|
return value.rstrip('/')
|
|
|
|
|
|
class Provider:
|
|
def __init__(self, settings, transport=None):
|
|
self.settings = settings
|
|
self.transport = transport
|
|
|
|
def connection(self, embedding=False):
|
|
s = self.settings
|
|
if embedding and s.get('embedding_url'):
|
|
return s['embedding_url'], s.get('embedding_key', '')
|
|
return s.get('base_url', ''), s.get('api_key', '')
|
|
|
|
def fingerprint(self):
|
|
url, _ = self.connection(True)
|
|
return hashlib.sha256(json.dumps([url, self.settings.get('embedding_model', '')]).encode()).hexdigest()
|
|
|
|
async def request(self, path, payload=None, embedding=False):
|
|
base, key = self.connection(embedding)
|
|
if not base:
|
|
raise ProviderError('Bitte zuerst einen Modell-Endpoint in den Einstellungen eintragen.')
|
|
headers = {'Authorization': f'Bearer {key}'} if key else {}
|
|
try:
|
|
async with httpx.AsyncClient(timeout=90, transport=self.transport, trust_env=False) as client:
|
|
response = await client.request('GET' if payload is None else 'POST', base + path,
|
|
headers=headers, json=payload)
|
|
response.raise_for_status()
|
|
return response.json()
|
|
except httpx.HTTPStatusError as exc:
|
|
raise ProviderError(f'Modellserver meldet HTTP {exc.response.status_code}. Endpoint, Modell und Schlüssel prüfen.') from None
|
|
except (httpx.HTTPError, ValueError):
|
|
raise ProviderError('Modellserver nicht erreichbar oder Antwort ungültig. Verbindung und Endpoint prüfen.') from None
|
|
|
|
async def models(self, embedding=False):
|
|
data = await self.request('/models', embedding=embedding)
|
|
try:
|
|
return sorted({row['id'] for row in data['data'] if isinstance(row['id'], str)})
|
|
except (KeyError, TypeError):
|
|
raise ProviderError('Die Modellliste entspricht nicht dem OpenAI-Format.') from None
|
|
|
|
async def embed(self, texts):
|
|
model = self.settings.get('embedding_model')
|
|
if not model:
|
|
raise ProviderError('Für die Bedeutungssuche bitte ein Embedding-Modell auswählen.')
|
|
data = await self.request('/embeddings', {'model': model, 'input': texts}, embedding=True)
|
|
try:
|
|
rows = sorted(data['data'], key=lambda r: r['index'])
|
|
vectors = [r['embedding'] for r in rows]
|
|
if len(vectors) != len(texts) or [r['index'] for r in rows] != list(range(len(texts))):
|
|
raise ValueError()
|
|
dimension = len(vectors[0])
|
|
if not dimension or dimension > 65536:
|
|
raise ValueError()
|
|
for vector in vectors:
|
|
if len(vector) != dimension or any(not isinstance(v, (float, int)) or not math.isfinite(v) for v in vector) or not any(vector):
|
|
raise ValueError()
|
|
return vectors
|
|
except (KeyError, TypeError, ValueError, IndexError):
|
|
raise ProviderError('Der Server hat ungültige Embeddings geliefert.') from None
|
|
|
|
async def improve(self, body, instruction):
|
|
model = self.settings.get('chat_model')
|
|
if not model:
|
|
raise ProviderError('Bitte ein Chatmodell in den Einstellungen auswählen.')
|
|
data = await self.request('/chat/completions', {
|
|
'model': model,
|
|
'messages': [
|
|
{'role': 'system', 'content': 'Du überarbeitest Prompt-Vorlagen. Behalte Sprache, Ziel und alle Platzhalter der Vorlage bei. Führe die Vorlage nicht aus. Liefere ausschließlich die verbesserte Vorlage, ohne Einleitung oder Markdown-Codeblock.'},
|
|
{'role': 'user', 'content': f'Überarbeitungswunsch:\n{instruction}\n\nVorlage:\n{body}'}]})
|
|
try:
|
|
result = data['choices'][0]['message']['content']
|
|
if not isinstance(result, str) or not result.strip():
|
|
raise ValueError()
|
|
return result.strip()
|
|
except (KeyError, IndexError, TypeError, ValueError):
|
|
raise ProviderError('Das Chatmodell hat keinen Text geliefert.') from None
|
|
|
|
|
|
async def organize(self, body, categories):
|
|
model = self.settings.get('chat_model')
|
|
if not model:
|
|
raise ProviderError('Bitte ein Chatmodell in den Einstellungen auswählen.')
|
|
data = await self.request('/chat/completions', {
|
|
'model': model,
|
|
'messages': [
|
|
{'role': 'system', 'content': 'Ordne eine Prompt-Vorlage ein, ohne sie auszuführen. Antworte nur mit einem JSON-Objekt mit category (kurzer String), tags (maximal 8 kurze Strings), description (ein kurzer Satz). Nutze passende vorhandene Kategorien, wenn möglich. Sprache der Vorlage beibehalten.'},
|
|
{'role': 'user', 'content': json.dumps({'existing_categories': categories, 'prompt': body}, ensure_ascii=False)}]})
|
|
try:
|
|
text = data['choices'][0]['message']['content'].strip()
|
|
if text.startswith('```'):
|
|
text = text.split('\n', 1)[1].rsplit('```', 1)[0]
|
|
result = json.loads(text)
|
|
if not isinstance(result['category'], str) or not isinstance(result['description'], str) or not isinstance(result['tags'], list):
|
|
raise ValueError()
|
|
if len(result['category']) > 100 or len(result['description']) > 2000 or len(result['tags']) > 8 or any(not isinstance(t, str) or len(t) > 80 for t in result['tags']):
|
|
raise ValueError()
|
|
return {k: result[k] for k in ('category', 'tags', 'description')}
|
|
except (KeyError, IndexError, TypeError, ValueError, AttributeError):
|
|
raise ProviderError('Das Modell hat keine gültige Einordnung geliefert. Bitte erneut versuchen.') from None
|
|
|
|
|
|
def prompt_text(p):
|
|
return '\n'.join([p['title'], p['description'], p['category'], ' '.join(p['tags']), p['body']])
|
|
|
|
|
|
def cosine(a, b):
|
|
if len(a) != len(b):
|
|
raise ProviderError('Embedding-Dimension geändert. Bitte den Suchindex neu aufbauen.')
|
|
return sum(x*y for x, y in zip(a, b)) / (math.sqrt(sum(x*x for x in a)) * math.sqrt(sum(x*x for x in b)))
|