230 lines
8.3 KiB
Python
230 lines
8.3 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
|
|
from typing import Optional
|
|
|
|
from app.config import settings
|
|
from app.services.llm_client import PROVIDERS, PROVIDER_MODELS
|
|
|
|
__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",
|
|
]
|
|
|
|
|
|
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 "",
|
|
}
|
|
try:
|
|
from app.db import get_conn
|
|
|
|
with get_conn() as conn:
|
|
row = conn.execute(
|
|
"SELECT provider, model, api_key, api_base 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]
|
|
return cfg
|
|
|
|
|
|
def set_llm_config(*, provider: str | None = None, model: str | None = None,
|
|
api_key: str | None = None, api_base: str | None = None) -> dict:
|
|
"""Upsert the runtime LLM config row (id=1) and return the new effective config."""
|
|
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:
|
|
cfg["api_key"] = api_key
|
|
if api_base is not None:
|
|
cfg["api_base"] = api_base
|
|
|
|
with get_conn() as conn:
|
|
conn.execute(
|
|
"""INSERT INTO llm_config (id, provider, model, api_key, api_base, 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,
|
|
updated_at=CURRENT_TIMESTAMP""",
|
|
(cfg["provider"], cfg["model"], cfg["api_key"], cfg["api_base"]),
|
|
)
|
|
conn.commit()
|
|
return cfg
|
|
|
|
|
|
def provider_info() -> list[dict]:
|
|
"""Providers list for the UI: known models + whether an API key is required."""
|
|
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": name.capitalize(),
|
|
"default_model": default_model,
|
|
"models": models,
|
|
"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."""
|
|
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"]),
|
|
}
|
|
|
|
|
|
def get_user_llm_key(user_id: int, provider: str) -> Optional[dict]:
|
|
"""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 upsert_user_llm_key(user_id: int, provider: str, *, api_key: str = "",
|
|
api_base: str = "", default_model: str = "",
|
|
models: Optional[list[str]] = None) -> dict:
|
|
"""Upsert a user's provider key. Empty api_key keeps the existing one
|
|
(allows saving model/base without re-typing the key)."""
|
|
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 "")
|
|
new_base = api_base if api_base else (existing.get("api_base", "") if existing else "")
|
|
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 []
|
|
)
|
|
with get_conn() as conn:
|
|
conn.execute(
|
|
"""INSERT INTO user_llm_keys (user_id, provider, api_key, api_base, default_model, models_json, 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,
|
|
updated_at=CURRENT_TIMESTAMP""",
|
|
(user_id, provider, new_key, new_base, new_model,
|
|
json.dumps(new_models, ensure_ascii=False)),
|
|
)
|
|
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),
|
|
}
|
|
|
|
|
|
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 use `GET {base}/models` with a Bearer token;
|
|
Anthropic uses `x-api-key` + `anthropic-version`; Gemini an `x-goog-api-key`.
|
|
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 provider == "google":
|
|
headers = {"x-goog-api-key": api_key}
|
|
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)
|
|
resp.raise_for_status()
|
|
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)
|
|
return out[:300] |