La clé fonctionne: /v1/models 200 + chat 200 (vérifié live). L'échec venait
du modèle: mistral-large-latest (défaut PROVIDERS + tête de liste) n'est pas
servi par les plans d'abonnement basiques → 403 code 1910 « tier_not_allowed »,
et le Test connection sélectionnait systématiquement ce premier candidat.
- PROVIDERS mistral default → mistral-small-latest (tous plans)
- PROVIDER_MODELS réordonnée: tous-plans d'abord (ordre du test de connexion)
- mistral ∈ _CHAT_VALIDATED_PROVIDERS: le fetch modèles ne garde que les
modèles réellement servis (46 listés → 25 utilisables avec la clé du compte)
- migration 36: default_model bloqué (large/pixtral-large) reset en base
Test live déployé: {ok:true, model:mistral-small-latest, reply:PONG,
verified:true} · tests/test_agent.py 53/53 · nouveau test d'ancrage
test_mistral_defaults_are_tier_safe
447 lines
18 KiB
Python
447 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.http_client import shared_client
|
|
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.
|
|
# Mistral : le /models liste des modèles hors du plan d'abonnement du compte
|
|
# (403 « tier_not_allowed » ex. mistral-large sur plan basique) — la sonde
|
|
# chat ne garde que ceux réellement servis par la CLÉ de l'utilisateur.
|
|
_CHAT_VALIDATED_PROVIDERS = frozenset({"nvidia", "mistral"})
|
|
|
|
# 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.
|
|
"""
|
|
|
|
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 shared_client(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
|
|
|
|
|
|
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 shared_client(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)]
|