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