Files
CasaDePrompt/atelier/provider.py
T

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)))