Un 4xx (ex. 403 Mistral) affichait seulement 'Client error 403 Forbidden'. Le message d'erreur inclut desormais le corps renvoye par le fournisseur (modele non autorise, region bloquee, etc.) pour le test de connexion et la recuperation des modeles.
445 lines
18 KiB
Python
445 lines
18 KiB
Python
"""FlowDeck — Runtime LLM configuration store (v4.10.2).
|
|
|
|
Precedence (per conversation):
|
|
1. explicit `provider`/`model` passed to the run endpoint,
|
|
2. the user's saved key for that provider (`user_llm_keys`),
|
|
3. the global `llm_config` row (id=1) — admin UI,
|
|
4. `settings.llm_*` (.env), default « offline mock ».
|
|
|
|
Rows are created lazily, so .env stays the default until saved from the UI.
|
|
"""
|
|
from __future__ import annotations
|
|
|
|
import json
|
|
import time
|
|
|
|
from app.config import settings
|
|
from app.services.llm_client import PROVIDER_LABELS, PROVIDER_MODELS, PROVIDERS
|
|
|
|
# Providers whose /v1/models lists far more entries than /v1/chat/completions
|
|
# actually serves. The fetched list is validated (name filter + live probe)
|
|
# before being exposed as "usable models". NVIDIA exposes all of its catalog
|
|
# (embeddings, rerank, image/video/audio gen…) many of which answer 404 on
|
|
# chat completions — the exact failure the user hit.
|
|
_CHAT_VALIDATED_PROVIDERS = frozenset({"nvidia"})
|
|
|
|
# Markers that identify clearly non-chat models (embeddings, rerank, media gen…).
|
|
_NON_CHAT_MARKERS = (
|
|
"embed", "bge-", "rerank", "retriev", "tts", "asr", "stt", "whisper", "speech",
|
|
"transcrib", "translate", "image", "video", "audio", "music", "sound", "dall",
|
|
"stable", "diffus", "flux", "sora", "veo", "midjourney", "clip", "segmentation",
|
|
"ocr", "inpainting", "depth", "motion", "sento-",
|
|
)
|
|
|
|
|
|
def _is_likely_chat(model_id: str) -> bool:
|
|
ml = model_id.lower()
|
|
return not any(m in ml for m in _NON_CHAT_MARKERS)
|
|
|
|
__all__ = [
|
|
"get_llm_config", "set_llm_config", "provider_info",
|
|
"get_user_llm_key", "list_user_llm_keys", "upsert_user_llm_key",
|
|
"delete_user_llm_key", "fetch_provider_models",
|
|
"mark_llm_config_verified", "mark_user_llm_key_verified",
|
|
]
|
|
|
|
|
|
def get_llm_config() -> dict:
|
|
"""Return the effective LLM config — DB overrides .env when present."""
|
|
cfg = {
|
|
"provider": settings.llm_provider or "offline",
|
|
"model": settings.llm_model or "gpt-4o",
|
|
"api_key": settings.llm_api_key or "",
|
|
"api_base": settings.llm_api_base or "",
|
|
"verified": 0,
|
|
"verified_model": "",
|
|
"verified_at": "",
|
|
"last_error": "",
|
|
}
|
|
try:
|
|
from app.db import get_conn
|
|
|
|
with get_conn() as conn:
|
|
row = conn.execute(
|
|
"SELECT provider, model, api_key, api_base, verified, verified_model, "
|
|
"verified_at, last_error FROM llm_config WHERE id=1"
|
|
).fetchone()
|
|
except Exception: # noqa: BLE001 — DB not ready yet → env defaults
|
|
return cfg
|
|
if row:
|
|
for key in ("provider", "model", "api_key", "api_base"):
|
|
if row[key]:
|
|
cfg[key] = row[key]
|
|
if row["provider"]:
|
|
cfg["verified"] = row["verified"] or 0
|
|
cfg["verified_model"] = row["verified_model"] or ""
|
|
cfg["verified_at"] = row["verified_at"] or ""
|
|
cfg["last_error"] = row["last_error"] or ""
|
|
return cfg
|
|
|
|
|
|
def set_llm_config(*, provider: str | None = None, model: str | None = None,
|
|
api_key: str | None = None, api_base: str | None = None,
|
|
clear_keys: bool = False) -> dict:
|
|
"""Upsert the runtime LLM config row (id=1) and return the new effective config.
|
|
|
|
An empty `api_key`/`api_base` keeps the stored value (so an admin can tweak
|
|
the model/base without re-typing the key). Changing the API key resets the
|
|
`verified` flag — the provider has to pass a connection test again.
|
|
"""
|
|
from app.db import get_conn
|
|
|
|
cfg = get_llm_config()
|
|
if provider is not None:
|
|
cfg["provider"] = provider
|
|
if model is not None:
|
|
cfg["model"] = model
|
|
if api_key is not None and api_key.strip():
|
|
key_changed = cfg.get("api_key") != api_key.strip()
|
|
cfg["api_key"] = api_key.strip()
|
|
if key_changed:
|
|
cfg["verified"] = 0
|
|
cfg["verified_model"] = ""
|
|
cfg["last_error"] = ""
|
|
if api_base is not None:
|
|
# Normalise so a base equal to the provider default is stored as "use
|
|
# default" — a later correction of PROVIDERS then applies automatically.
|
|
cfg["api_base"] = _normalize_api_base(
|
|
provider or cfg.get("provider") or "offline", api_base
|
|
)
|
|
if clear_keys:
|
|
cfg["api_key"] = ""
|
|
cfg["api_base"] = ""
|
|
cfg["verified"] = 0
|
|
cfg["verified_model"] = ""
|
|
cfg["last_error"] = ""
|
|
|
|
with get_conn() as conn:
|
|
conn.execute(
|
|
"""INSERT INTO llm_config (id, provider, model, api_key, api_base,
|
|
verified, verified_model, last_error, updated_at)
|
|
VALUES (1, ?, ?, ?, ?, ?, ?, ?, CURRENT_TIMESTAMP)
|
|
ON CONFLICT(id) DO UPDATE SET
|
|
provider=excluded.provider, model=excluded.model,
|
|
api_key=excluded.api_key, api_base=excluded.api_base,
|
|
verified=excluded.verified, verified_model=excluded.verified_model,
|
|
last_error=excluded.last_error,
|
|
updated_at=CURRENT_TIMESTAMP""",
|
|
(cfg["provider"], cfg["model"], cfg["api_key"], cfg["api_base"],
|
|
cfg.get("verified", 0), cfg.get("verified_model", ""),
|
|
cfg.get("last_error", "")),
|
|
)
|
|
conn.commit()
|
|
return cfg
|
|
|
|
|
|
def mark_llm_config_verified(ok: bool, *, model: str = "", error: str = "") -> None:
|
|
"""Record the outcome of the admin 'Test connection' for the global default."""
|
|
from app.db import get_conn
|
|
|
|
with get_conn() as conn:
|
|
row = conn.execute("SELECT provider FROM llm_config WHERE id=1").fetchone()
|
|
provider = row["provider"] if row else (settings.llm_provider or "offline")
|
|
conn.execute(
|
|
"""INSERT INTO llm_config (id, provider, model, verified, verified_model,
|
|
verified_at, last_error, updated_at)
|
|
VALUES (1, ?, '', ?, ?, CASE WHEN ? THEN CURRENT_TIMESTAMP END, ?, CURRENT_TIMESTAMP)
|
|
ON CONFLICT(id) DO UPDATE SET
|
|
verified=excluded.verified,
|
|
verified_model=excluded.verified_model,
|
|
verified_at=excluded.verified_at,
|
|
last_error=excluded.last_error,
|
|
updated_at=CURRENT_TIMESTAMP""",
|
|
(provider, 1 if ok else 0, model if ok else "",
|
|
1 if ok else 0, "" if ok else error),
|
|
)
|
|
conn.commit()
|
|
|
|
|
|
def provider_info() -> list[dict]:
|
|
"""Providers list for the UI: known models + whether an API key is required.
|
|
|
|
``base_url`` is the endpoint the application uses by default for this
|
|
provider's chat requests (shown in the Settings "URL API" fields).
|
|
"""
|
|
out: list[dict] = []
|
|
for name, (base, default_model) in PROVIDERS.items():
|
|
models = list(PROVIDER_MODELS.get(name) or ())
|
|
if default_model and default_model not in models:
|
|
models.insert(0, default_model)
|
|
out.append({
|
|
"id": name,
|
|
"name": PROVIDER_LABELS.get(name) or name.replace("_", " ").title(),
|
|
"default_model": default_model,
|
|
"models": models,
|
|
"base_url": (base or ""),
|
|
"requires_key": name not in ("offline", "ollama"),
|
|
})
|
|
return out
|
|
|
|
|
|
# ── Per-user provider keys ──
|
|
|
|
|
|
def _mask_key(row) -> dict:
|
|
"""Public view of a stored key row: never exposes the raw api_key."""
|
|
def _get(name, default=""):
|
|
try:
|
|
v = row[name]
|
|
return default if v is None else v
|
|
except (KeyError, IndexError):
|
|
return default
|
|
return {
|
|
"provider": row["provider"],
|
|
"api_base": row["api_base"] or "",
|
|
"default_model": row["default_model"] or "",
|
|
"models": json.loads(row["models_json"] or "[]"),
|
|
"has_key": bool(row["api_key"]),
|
|
"verified": bool(_get("verified", 0)),
|
|
"verified_model": _get("verified_model") or "",
|
|
"last_error": _get("last_error") or "",
|
|
}
|
|
|
|
|
|
def get_user_llm_key(user_id: int, provider: str) -> dict | None:
|
|
"""Return a stored key row (includes the raw api_key — server-side only)."""
|
|
from app.db import get_conn
|
|
|
|
provider = provider.lower()
|
|
with get_conn() as conn:
|
|
row = conn.execute(
|
|
"SELECT * FROM user_llm_keys WHERE user_id=? AND provider=?",
|
|
(user_id, provider),
|
|
).fetchone()
|
|
return dict(row) if row else None
|
|
|
|
|
|
def list_user_llm_keys(user_id: int) -> list[dict]:
|
|
"""Public (masked) list of the user's saved provider keys."""
|
|
from app.db import get_conn
|
|
|
|
with get_conn() as conn:
|
|
rows = conn.execute(
|
|
"SELECT * FROM user_llm_keys WHERE user_id=? ORDER BY provider",
|
|
(user_id,),
|
|
).fetchall()
|
|
return [_mask_key(r) for r in rows]
|
|
|
|
|
|
def _default_api_base(provider: str) -> str:
|
|
"""The OpenAI-compatible base URL the app uses by default for a provider."""
|
|
base = (PROVIDERS.get(provider.lower()) or (None, None))[0]
|
|
return (base or "").rstrip("/")
|
|
|
|
|
|
def _normalize_api_base(provider: str, api_base: str) -> str:
|
|
"""Drop an api_base that just repeats the provider default.
|
|
|
|
Storing the default as a per-user override freezes it: a later correction
|
|
of the provider URL (e.g. Cohere `/v2` → `/compatibility/v1`) would never
|
|
apply. An empty value means "use the provider default".
|
|
"""
|
|
value = (api_base or "").strip()
|
|
if value.rstrip("/") == _default_api_base(provider):
|
|
return ""
|
|
return value
|
|
|
|
|
|
def upsert_user_llm_key(user_id: int, provider: str, *, api_key: str = "",
|
|
api_base: str | None = None, default_model: str = "",
|
|
models: list[str] | None = None) -> dict:
|
|
"""Upsert a user's provider key. Empty api_key keeps the existing one
|
|
(allows saving model/base without re-typing the key). Saving a *different*
|
|
key resets the `verified` flag so the provider must pass a test again.
|
|
|
|
``api_base`` uses ``None`` to mean "keep the stored value" and an empty
|
|
string to explicitly reset it to the provider default.
|
|
"""
|
|
from app.db import get_conn
|
|
|
|
provider = provider.lower()
|
|
existing = get_user_llm_key(user_id, provider)
|
|
new_key = api_key if api_key else (existing.get("api_key", "") if existing else "")
|
|
if api_base is None:
|
|
new_base = (existing.get("api_base", "") if existing else "")
|
|
else:
|
|
new_base = _normalize_api_base(provider, api_base)
|
|
new_model = default_model if default_model else (existing.get("default_model", "") if existing else "")
|
|
new_models = models if models is not None else (
|
|
json.loads(existing["models_json"]) if existing and existing.get("models_json") else []
|
|
)
|
|
key_changed = bool(api_key) and (not existing or existing.get("api_key", "") != api_key)
|
|
# A changed key invalidates the previous verification; keep it otherwise.
|
|
verified = 0 if key_changed else (existing.get("verified") or 0) if existing else 0
|
|
with get_conn() as conn:
|
|
conn.execute(
|
|
"""INSERT INTO user_llm_keys (user_id, provider, api_key, api_base, default_model, models_json, verified, updated_at)
|
|
VALUES (?, ?, ?, ?, ?, ?, ?, CURRENT_TIMESTAMP)
|
|
ON CONFLICT(user_id, provider) DO UPDATE SET
|
|
api_key=excluded.api_key, api_base=excluded.api_base,
|
|
default_model=excluded.default_model, models_json=excluded.models_json,
|
|
verified=excluded.verified,
|
|
updated_at=CURRENT_TIMESTAMP""",
|
|
(user_id, provider, new_key, new_base, new_model,
|
|
json.dumps(new_models, ensure_ascii=False), verified),
|
|
)
|
|
conn.commit()
|
|
return get_user_llm_key(user_id, provider) or {
|
|
"provider": provider, "api_key": new_key, "api_base": new_base,
|
|
"default_model": new_model, "models_json": json.dumps(new_models, ensure_ascii=False),
|
|
"verified": verified,
|
|
}
|
|
|
|
|
|
def mark_user_llm_key_verified(user_id: int, provider: str, ok: bool, *,
|
|
model: str = "", error: str = "") -> None:
|
|
"""Record the outcome of a connection test on one of the user's providers."""
|
|
from app.db import get_conn
|
|
|
|
provider = provider.lower()
|
|
with get_conn() as conn:
|
|
conn.execute(
|
|
"""UPDATE user_llm_keys
|
|
SET verified=?, verified_model=?, last_error=?,
|
|
verified_at=CASE WHEN ? THEN CURRENT_TIMESTAMP END,
|
|
updated_at=CURRENT_TIMESTAMP
|
|
WHERE user_id=? AND provider=?""",
|
|
(1 if ok else 0, model if ok else "", "" if ok else error,
|
|
1 if ok else 0, user_id, provider),
|
|
)
|
|
conn.commit()
|
|
|
|
|
|
def delete_user_llm_key(user_id: int, provider: str) -> None:
|
|
from app.db import get_conn
|
|
|
|
provider = provider.lower()
|
|
with get_conn() as conn:
|
|
conn.execute(
|
|
"DELETE FROM user_llm_keys WHERE user_id=? AND provider=?",
|
|
(user_id, provider),
|
|
)
|
|
conn.commit()
|
|
|
|
|
|
async def fetch_provider_models(provider: str, *, api_key: str = "",
|
|
api_base: str = "", timeout: int = 20) -> list[str]:
|
|
"""Fetch the live model list from a provider (best-effort, no mock).
|
|
|
|
OpenAI-compatible providers (including Google's `/openai` surface and
|
|
Cohere's compatibility API) use `GET {base}/models` with a Bearer token;
|
|
Anthropic's native model listing uses `x-api-key` + `anthropic-version`.
|
|
Returns a de-duplicated list capped at 300 models.
|
|
"""
|
|
import httpx
|
|
|
|
provider = provider.lower()
|
|
base = (api_base or "").strip() or (PROVIDERS.get(provider) or (None, None))[0]
|
|
if not base:
|
|
return [] # offline — nothing to fetch
|
|
|
|
base_url = base.rstrip("/")
|
|
headers: dict = {}
|
|
if provider == "anthropic":
|
|
headers = {"x-api-key": api_key, "anthropic-version": "2023-06-01"}
|
|
elif api_key:
|
|
headers = {"Authorization": f"Bearer {api_key}"}
|
|
|
|
async with httpx.AsyncClient(timeout=timeout) as client:
|
|
resp = await client.get(f"{base_url}/models", headers=headers)
|
|
if resp.status_code >= 400:
|
|
body = (resp.text or "").strip()
|
|
if len(body) > 500:
|
|
body = body[:500] + "…"
|
|
raise RuntimeError(
|
|
f"{resp.status_code} {resp.reason_phrase} ({resp.url}): {body}"
|
|
)
|
|
data = resp.json()
|
|
|
|
ids: list[str] = []
|
|
for item in data.get("data") or []:
|
|
if not isinstance(item, dict):
|
|
continue
|
|
i = (item.get("id") or "").strip()
|
|
if i:
|
|
ids.append(i)
|
|
for item in data.get("models") or []:
|
|
if not isinstance(item, dict):
|
|
continue
|
|
n = (item.get("name") or item.get("id") or "").strip()
|
|
if provider == "google" and n.startswith("models/"):
|
|
n = n[len("models/"):]
|
|
if n:
|
|
ids.append(n)
|
|
seen: set[str] = set()
|
|
out: list[str] = []
|
|
for i in ids:
|
|
if i not in seen:
|
|
seen.add(i)
|
|
out.append(i)
|
|
|
|
# Providers like NVIDIA list their whole catalog, most of which is NOT served
|
|
# by /v1/chat/completions (404 sur « model not found »). Filter by name first,
|
|
# then probe the survivors with a minimal chat call so only usable models stay.
|
|
if provider in _CHAT_VALIDATED_PROVIDERS and out:
|
|
candidates = [m for m in out if _is_likely_chat(m)] or out
|
|
validated = await _validate_chat_models(base_url, api_key, candidates)
|
|
# Ne vidons jamais la liste : en cas d'échec de validation (débit limité,
|
|
# indisponibilité passagère) on garde la liste filtrée par nom.
|
|
out = validated if validated else candidates
|
|
|
|
return out[:300]
|
|
|
|
|
|
async def _validate_chat_models(base_url: str, api_key: str, candidates: list[str],
|
|
*, timeout: float = 12.0, concurrency: int = 5,
|
|
deadline: float = 90.0, attempts: int = 2) -> list[str]:
|
|
"""Probe `POST {base_url}/chat/completions` for each candidate and keep only
|
|
the models that genuinely answer — using the exact payload the app sends at
|
|
runtime (`temperature: 0.2`, no `max_tokens`).
|
|
|
|
Hard failures (404 model inconnu, 410 modèle retiré…) exclude the model.
|
|
A 429 (rate limit) is kept: it proves the route exists, so the model is
|
|
usable once the quota frees up. Transient failures (timeouts, 5xx, network)
|
|
are retried once before giving up, so slow-but-working models survive.
|
|
"""
|
|
import asyncio
|
|
|
|
import httpx
|
|
|
|
sem = asyncio.Semaphore(concurrency)
|
|
start = time.monotonic()
|
|
headers = {"Content-Type": "application/json"}
|
|
if api_key:
|
|
headers["Authorization"] = f"Bearer {api_key}"
|
|
url = f"{base_url.rstrip('/')}/chat/completions"
|
|
payload_tpl = {
|
|
"messages": [{"role": "user", "content": "ping"}],
|
|
"temperature": 0.2,
|
|
}
|
|
|
|
async def probe(model: str) -> str | None:
|
|
if time.monotonic() - start > deadline:
|
|
return None
|
|
payload: dict = {"model": model, **payload_tpl}
|
|
for attempt in range(attempts):
|
|
if time.monotonic() - start > deadline:
|
|
return None
|
|
try:
|
|
async with sem:
|
|
async with httpx.AsyncClient(timeout=timeout) as client:
|
|
resp = await client.post(url, headers=headers, json=payload)
|
|
code = resp.status_code
|
|
if code < 300 or code == 429:
|
|
return model
|
|
if code < 500: # 400/401/403/404/410… → définitivement non utilisable
|
|
return None
|
|
# 5xx → passage passager, on retente une fois
|
|
except Exception: # noqa: BLE001 — timeouts / connexion → on retente
|
|
if attempt >= attempts - 1:
|
|
return None
|
|
return None
|
|
|
|
results = await asyncio.gather(*(probe(m) for m in candidates))
|
|
return [m for m in results if isinstance(m, str)]
|