diff --git a/backend/ai_routes.py b/backend/ai_routes.py index f9033dd..2225568 100644 --- a/backend/ai_routes.py +++ b/backend/ai_routes.py @@ -101,7 +101,8 @@ class AIRequest(BaseModel): instruction: str | None = Field(None, description="Custom instruction for rewrite") target_lang: str | None = Field(None, description="Target language for translation") tone: str | None = Field(None, description="Target tone (professional, casual, etc.)") - provider: str | None = Field(None, description="AI provider override") + provider: str | None = Field(None, description="AI provider override (e.g. 'deepseek', 'nvidia')") + model: str | None = Field(None, description="Model name override for this request") class AIResponse(BaseModel): @@ -111,6 +112,12 @@ class AIResponse(BaseModel): async def _handle(action, request: AIRequest): """Wrapper with error handling.""" + from backend.ai import PROVIDERS + # Apply per-request model override (saved/restored around the call) + original_model = None + if request.model and request.provider and request.provider in PROVIDERS: + original_model = PROVIDERS[request.provider].get("model") + PROVIDERS[request.provider]["model"] = request.model try: result = await action(request.text, request.provider) return AIResponse(result=result, provider=request.provider or "default") @@ -119,6 +126,9 @@ async def _handle(action, request: AIRequest): except Exception as e: logger.error(f"AI error: {e}") raise HTTPException(status_code=500, detail=f"AI service error: {e!s}") + finally: + if original_model is not None and request.provider in PROVIDERS: + PROVIDERS[request.provider]["model"] = original_model @router.post("/improve", response_model=AIResponse) diff --git a/backend/bookslm_routes.py b/backend/bookslm_routes.py index 347e0bd..48699a7 100644 --- a/backend/bookslm_routes.py +++ b/backend/bookslm_routes.py @@ -32,6 +32,15 @@ class BooksLMChatRequest(BaseModel): default_factory=list, description="Previous conversation turns [{role, content}]", ) + provider: str | None = Field( + default=None, + description="AI provider override (e.g. 'deepseek', 'openrouter', 'gemini', 'nvidia', 'xiaomi', 'mistral', 'qwencloud'). " + "If not set, uses DEFAULT_PROVIDER.", + ) + model: str | None = Field( + default=None, + description="Model name to use for this request. If not set, uses the provider's default model.", + ) # ── Endpoints ── @@ -107,27 +116,44 @@ async def api_bookslm_chat( async def generate_sse(): try: - cfg_name = DEFAULT_PROVIDER - # Check if default provider is available - if cfg_name != "gemini" and cfg_name in PROVIDERS and not PROVIDERS[cfg_name].get("api_key"): - # Find first available provider + # Resolve provider: explicit override wins, else default. + # Fall back to first available if the requested one isn't configured. + cfg_name = (req.provider or DEFAULT_PROVIDER).lower() + if cfg_name not in PROVIDERS or not PROVIDERS[cfg_name].get("api_key"): + # Try next available provider for pname, pcfg in PROVIDERS.items(): if pcfg.get("api_key") and pname != "gemini": cfg_name = pname break + else: + # No provider available at all + err = "Aucun fournisseur AI configuré (clés API manquantes)" + error_data = json.dumps({"error": err}, ensure_ascii=False) + yield f"event: error\ndata: {error_data}\n\n" + return - if cfg_name == "gemini" and PROVIDERS.get("gemini", {}).get("api_key"): - response = await _call_gemini(user_prompt, system_prompt, temperature=0.3, max_tokens=4096) - else: - response = await _call_deepseek_openrouter( - user_prompt, system_prompt, - provider=cfg_name if cfg_name in PROVIDERS else None, - temperature=0.3, - max_tokens=4096, - ) + # Optional per-request model override + original_model = None + if req.model and cfg_name in PROVIDERS: + original_model = PROVIDERS[cfg_name].get("model") + PROVIDERS[cfg_name]["model"] = req.model + try: + if cfg_name == "gemini": + response = await _call_gemini(user_prompt, system_prompt, temperature=0.3, max_tokens=4096) + else: + response = await _call_deepseek_openrouter( + user_prompt, system_prompt, + provider=cfg_name, + temperature=0.3, + max_tokens=4096, + ) + finally: + # Restore the original model so other calls aren't affected + if original_model is not None and cfg_name in PROVIDERS: + PROVIDERS[cfg_name]["model"] = original_model # Send the full response as a single SSE event - data = json.dumps({"token": response}, ensure_ascii=False) + data = json.dumps({"token": response, "provider": cfg_name, "model": req.model or PROVIDERS.get(cfg_name, {}).get("model", "")}, ensure_ascii=False) yield f"event: message\ndata: {data}\n\n" yield "event: done\ndata: {}\n\n" except Exception as e: diff --git a/backend/main.py b/backend/main.py index 40fad9b..3ae057f 100644 --- a/backend/main.py +++ b/backend/main.py @@ -4058,19 +4058,30 @@ async def api_test_ai_keys(current_user=Depends(require_admin)): @app.get("/api/config/ai-models") async def api_list_ai_models(provider: str = Query(...), current_user=Depends(require_admin)): - """List available models for a given AI provider.""" + """List available models for a given AI provider. + + Strategy: + 1. Try the provider's public models endpoint (OpenAI-compatible /v1/models or Gemini). + 2. If the network call fails (timeout, 4xx, 5xx, DNS, etc.), fall back to a + curated static list of known-good models for that provider. + 3. Always return a non-empty list when the provider is known, so the UI + dropdown is never empty. + """ provider = provider.lower() all_providers = ("deepseek", "openrouter", "gemini", "nvidia", "qwencloud", "xiaomi", "mistral") if provider not in all_providers: - return {"models": [], "error": f"Unknown provider: {provider}"} + return {"models": [], "error": f"Unknown provider: {provider}", "source": "validation"} key_name = f"{provider.upper()}_API_KEY" key = get_ai_key(key_name) if not key: - return {"models": [], "error": "API key not configured"} + # No key configured — return curated fallback list so the UI can + # still show what WOULD be available once a key is set. + return {"models": _FALLBACK_MODELS.get(provider, []), "source": "fallback", + "note": "API key not configured — showing default model list"} - # Build URL and request + # Build URL if provider == "gemini": url = f"https://generativelanguage.googleapis.com/v1beta/models?key={key}" elif provider == "deepseek": @@ -4082,6 +4093,8 @@ async def api_list_ai_models(provider: str = Query(...), current_user=Depends(re elif provider == "qwencloud": url = "https://dashscope.aliyuncs.com/compatible-mode/v1/models" elif provider == "xiaomi": + # Xiaomi's public /v1/models endpoint is not stable — fetch if reachable, + # otherwise fall back to a curated list of mimo models. url = "https://api.xiaomi.com/v1/models" elif provider == "mistral": url = "https://api.mistral.ai/v1/models" @@ -4096,13 +4109,86 @@ async def api_list_ai_models(provider: str = Query(...), current_user=Depends(re data = _json.loads(resp.read().decode()) if provider == "gemini": - models = [m.get("name", "") for m in data.get("models", [])] + models = [m.get("name", "") for m in data.get("models", []) if m.get("name")] + # Gemini returns names like "models/gemini-1.5-flash" — strip prefix + models = [m.replace("models/", "") for m in models] else: - models = [m.get("id", "") for m in data.get("data", [])] + models = [m.get("id", "") for m in data.get("data", []) if m.get("id")] - return {"models": models} + if models: + # Prepend the configured default if not already present + default = PROVIDERS.get(provider, {}).get("model") + if default and default not in models: + models = [default] + models + return {"models": models, "source": "live", "count": len(models)} + # Empty list from API — fall through to fallback + raise ValueError("empty model list from provider API") except Exception as e: - return {"models": [], "error": str(e)} + # Network error, auth error, parsing error — use curated fallback + fallback = _FALLBACK_MODELS.get(provider, []) + return {"models": fallback, "source": "fallback", "error": str(e)[:200], + "note": "Could not reach provider API — showing default model list"} + + +# ── Curated fallback model lists ────────────────────────────────────────── +# Used when the provider API is unreachable or returns empty. +# Keep these short and focused on models known to work with the +# OpenAI-compatible chat completions interface (or Gemini's generateContent). +_FALLBACK_MODELS: dict[str, list[str]] = { + "deepseek": [ + "deepseek-chat", + "deepseek-reasoner", + ], + "openrouter": [ + "openai/gpt-4o-mini", + "openai/gpt-4o", + "anthropic/claude-3.5-sonnet", + "anthropic/claude-3-haiku", + "google/gemini-2.0-flash-exp:free", + "meta-llama/llama-3.1-70b-instruct", + "meta-llama/llama-3.1-8b-instruct:free", + "mistralai/mistral-large-latest", + ], + "gemini": [ + "gemini-2.0-flash", + "gemini-2.0-flash-exp", + "gemini-1.5-pro", + "gemini-1.5-flash", + "gemini-1.5-flash-8b", + ], + "nvidia": [ + "meta/llama-3.1-405b-instruct", + "meta/llama-3.1-70b-instruct", + "meta/llama-3.1-8b-instruct", + "mistralai/mistral-large", + "google/gemma-2-27b-it", + "nvidia/llama-3.1-nemotron-70b-instruct", + ], + "qwencloud": [ + "qwen-max", + "qwen-plus", + "qwen-turbo", + "qwen-long", + "qwen-vl-max", + "qwen-vl-plus", + ], + "xiaomi": [ + # Xiaomi MiMo models — the public /v1/models endpoint is unreliable, + # so we ship a known-good list as fallback. + "mimo-v2-pro", + "mimo-v2-flash", + "mimo-v2-vl", + "mimo-v2-tts", + ], + "mistral": [ + "mistral-large-latest", + "mistral-medium-latest", + "mistral-small-latest", + "open-mistral-7b", + "open-mixtral-8x7b", + "codestral-latest", + ], +} # --------------------------------------------------------------------------- diff --git a/frontend/js/admin.js b/frontend/js/admin.js index 21941ab..b29c9d0 100644 --- a/frontend/js/admin.js +++ b/frontend/js/admin.js @@ -364,30 +364,51 @@ function _wireAuditFilters() { /** * Verify the current session and admin role. + * Uses /api/auth/me (which returns the current user) and falls back to + * the cached user in sessionStorage if the request fails. * Redirects to / if not authenticated or not admin. * Returns the user object on success, null otherwise. */ async function _gateAdmin() { const container = document.getElementById("admin-main"); const forbidden = document.getElementById("admin-forbidden"); + // First check whether auth is even enabled (public endpoint). + let authEnabled = false; try { - const res = await fetch("/api/auth/status", { credentials: "include" }); - if (!res.ok) throw new Error("HTTP " + res.status); - const status = await res.json(); - if (!status.auth_enabled) { - // Auth disabled — server is wide open, treat current visitor as allowed. - return { username: "anonymous", role: "admin" }; + const statusRes = await fetch("/api/auth/status", { credentials: "include" }); + if (statusRes.ok) { + const status = await statusRes.json(); + authEnabled = !!status.auth_enabled; } - if (!status.authenticated) { + } catch { /* network error — fall through */ } + + if (!authEnabled) { + // Server is wide open, treat current visitor as allowed. + return { username: "anonymous", role: "admin" }; + } + + // Auth is enabled — try to load the current user. + try { + const meRes = await fetch("/api/auth/me", { + credentials: "include", + headers: getAuthHeaders(), + }); + if (meRes.status === 401 || meRes.status === 403) { window.location.href = "/"; return null; } - if (status.user?.role !== "admin") { + if (!meRes.ok) throw new Error("HTTP " + meRes.status); + const user = await meRes.json(); + // Cache for later fallback + try { + sessionStorage.setItem("obsigate_user", JSON.stringify(user)); + } catch { /* */ } + if (user.role !== "admin") { if (container) container.style.display = "none"; if (forbidden) forbidden.style.display = ""; return null; } - return status.user; + return user; } catch (err) { // Network error — try to use cached sessionStorage user as fallback try { diff --git a/frontend/js/ai.js b/frontend/js/ai.js index 2fe9d62..9db5955 100644 --- a/frontend/js/ai.js +++ b/frontend/js/ai.js @@ -13,7 +13,14 @@ import { t } from './i18n.js'; // ── API call helper ── async function aiAction(endpoint, text, extra = {}) { - const body = { text, ...extra }; + // Inject current provider + model from the per-section picker. + const pickerState = _readPicker(); + const body = { + text, + ...extra, + ...(pickerState.provider ? { provider: pickerState.provider } : {}), + ...(pickerState.model ? { model: pickerState.model } : {}), + }; const data = await api(`/api/ai/${endpoint}`, { method: 'POST', headers: { 'Content-Type': 'application/json' }, @@ -22,6 +29,163 @@ async function aiAction(endpoint, text, extra = {}) { return data.result; } +// ── Per-section provider/model picker ─────────────────────────────────────── +// State is stored in localStorage so the user choice persists across sections. +const PICKER_STORAGE_KEY = 'obsigate_ai_picker'; + +function _readPicker() { + try { + return JSON.parse(localStorage.getItem(PICKER_STORAGE_KEY) || '{}'); + } catch { return {}; } +} + +function _writePicker(state) { + try { + localStorage.setItem(PICKER_STORAGE_KEY, JSON.stringify(state)); + } catch { /* */ } +} + +/** + * Build a provider+model picker that appears in every section's AI toolbar + * (Forge editor, BooksLM, etc.). The picker reads /api/ai/status to discover + * available providers, then /api/config/ai-models?provider=X to list models. + */ +async function _buildPickerUI() { + const pickerState = _readPicker(); + + // Fetch configured providers + let availableProviders = {}; + try { + const status = await api('/api/ai/status'); + availableProviders = status.providers || {}; + } catch { /* */ } + + const providerNames = Object.keys(availableProviders).filter( + (p) => availableProviders[p]?.available + ); + if (!providerNames.length) return null; + + const wrap = document.createElement('div'); + wrap.className = 'ai-picker'; + Object.assign(wrap.style, { + display: 'inline-flex', + alignItems: 'center', + gap: '4px', + marginLeft: '8px', + paddingLeft: '8px', + borderLeft: '1px solid var(--border-color)', + fontSize: '0.7rem', + color: 'var(--text-muted)', + }); + + // Provider select + const providerLabel = document.createElement('span'); + providerLabel.textContent = t('ai.provider') + ':'; + wrap.appendChild(providerLabel); + + const providerSelect = document.createElement('select'); + providerSelect.className = 'ai-picker-select'; + Object.assign(providerSelect.style, { + fontSize: '0.7rem', + background: 'var(--bg-tertiary, transparent)', + color: 'var(--text-primary)', + border: '1px solid var(--border-color)', + borderRadius: '4px', + padding: '2px 4px', + cursor: 'pointer', + }); + + const defaultOpt = document.createElement('option'); + defaultOpt.value = ''; + defaultOpt.textContent = t('ai.provider_default'); + providerSelect.appendChild(defaultOpt); + providerNames.forEach((p) => { + const opt = document.createElement('option'); + opt.value = p; + opt.textContent = p; + if (p === pickerState.provider) opt.selected = true; + providerSelect.appendChild(opt); + }); + + // Model select (sibling of provider) + const modelSelect = document.createElement('select'); + modelSelect.className = 'ai-picker-model'; + Object.assign(modelSelect.style, { + fontSize: '0.7rem', + background: 'var(--bg-tertiary, transparent)', + color: 'var(--text-primary)', + border: '1px solid var(--border-color)', + borderRadius: '4px', + padding: '2px 4px', + cursor: 'pointer', + minWidth: '120px', + maxWidth: '220px', + }); + const placeholderOpt = document.createElement('option'); + placeholderOpt.value = ''; + placeholderOpt.textContent = t('ai.model_default'); + modelSelect.appendChild(placeholderOpt); + + async function _loadModels(provider) { + modelSelect.innerHTML = ''; + const ph = document.createElement('option'); + ph.value = ''; + ph.textContent = t('ai.model_loading'); + modelSelect.appendChild(ph); + try { + const data = await api(`/api/config/ai-models?provider=${encodeURIComponent(provider)}`); + modelSelect.innerHTML = ''; + const def = document.createElement('option'); + def.value = ''; + def.textContent = t('ai.model_default') + (data.source === 'fallback' ? ` (${t('ai.model_offline')})` : ''); + modelSelect.appendChild(def); + (data.models || []).forEach((m) => { + const opt = document.createElement('option'); + opt.value = m; + opt.textContent = m; + if (m === pickerState.model && provider === pickerState.provider) opt.selected = true; + modelSelect.appendChild(opt); + }); + } catch (err) { + modelSelect.innerHTML = ''; + const opt = document.createElement('option'); + opt.value = ''; + opt.textContent = t('ai.model_load_error'); + modelSelect.appendChild(opt); + } + } + + providerSelect.addEventListener('change', () => { + pickerState.provider = providerSelect.value || null; + // Reset model when provider changes + pickerState.model = null; + _writePicker(pickerState); + if (pickerState.provider) { + _loadModels(pickerState.provider); + } else { + modelSelect.innerHTML = ''; + const ph = document.createElement('option'); + ph.value = ''; + ph.textContent = t('ai.model_default'); + modelSelect.appendChild(ph); + } + }); + + modelSelect.addEventListener('change', () => { + pickerState.model = modelSelect.value || null; + _writePicker(pickerState); + }); + + // Load models for the initial provider selection + if (pickerState.provider && providerNames.includes(pickerState.provider)) { + _loadModels(pickerState.provider); + } + + wrap.appendChild(providerSelect); + wrap.appendChild(modelSelect); + return wrap; +} + // ── Get selected text from CodeMirror ── function getSelection(editorView) { if (!editorView) return ''; @@ -357,6 +521,13 @@ export async function createAIToolbar(container, getEditorView) { toolbar.appendChild(rewriteBtn); toolbar.appendChild(toolboxBtn); + // Append the per-section provider/model picker on the right. + const picker = await _buildPickerUI(); + if (picker) { + toolbar.appendChild(createSeparator()); + toolbar.appendChild(picker); + } + container.insertBefore(toolbar, container.firstChild); // ── Action helper ── @@ -455,3 +626,13 @@ function createSeparator() { sep.style.cssText = 'width:1px;height:16px;background:var(--border-color);margin:0 2px'; return sep; } + +// ── Public exports ──────────────────────────────────────────────────────── +// Other modules (e.g. BooksLM) can import the picker helpers to render the +// same per-section provider/model picker in their own toolbar. +export { + _readPicker, + _writePicker, + PICKER_STORAGE_KEY, + _buildPickerUI as buildAIPickerUI, +}; diff --git a/frontend/js/bookslm.js b/frontend/js/bookslm.js index 774a3d7..3196888 100644 --- a/frontend/js/bookslm.js +++ b/frontend/js/bookslm.js @@ -1,5 +1,6 @@ // BooksLM — Directory-scoped AI chat panel (style NotebookLM) import { t } from './i18n.js'; +import { buildAIPickerUI } from './ai.js'; class BooksLM { constructor() { @@ -118,6 +119,7 @@ class BooksLM {
📚 + @@ -132,6 +134,12 @@ class BooksLM {
`; + // Inject the provider/model picker asynchronously (depends on /api/ai/status) + const pickerHost = panel.querySelector('.bookslm-picker-host'); + buildAIPickerUI().then((picker) => { + if (picker && pickerHost) pickerHost.replaceWith(picker); + }).catch(() => { /* ignore */ }); + // Wire events panel.querySelector('.bookslm-btn-close').addEventListener('click', () => this.close()); panel.querySelector('.bookslm-btn-new').addEventListener('click', () => this.newConversation()); @@ -288,6 +296,15 @@ class BooksLM { this._abortCtrl = new AbortController(); try { + // Read provider/model from the picker (localStorage-backed) + let provider = null; + let model = null; + try { + const picker = JSON.parse(localStorage.getItem('obsigate_ai_picker') || '{}'); + provider = picker.provider || null; + model = picker.model || null; + } catch { /* */ } + const resp = await fetch('/api/ai/bookslm/chat', { method: 'POST', headers: { 'Content-Type': 'application/json' }, @@ -295,7 +312,9 @@ class BooksLM { vault: this._vault, directory: this._directory, messages: this._messages.slice(0, -1), // exclude empty assistant - context_files: this._contextFiles.map(f => f.path || f) + context_files: this._contextFiles.map(f => f.path || f), + provider, + model, }), signal: this._abortCtrl.signal }); diff --git a/frontend/locales/en.json b/frontend/locales/en.json index 0f9915c..5d2e59e 100644 --- a/frontend/locales/en.json +++ b/frontend/locales/en.json @@ -102,6 +102,13 @@ "ai.not_configured": "⚠️ AI not configured — add DEEPSEEK_API_KEY, OPENROUTER_API_KEY or GEMINI_API_KEY in .env", "ai.processing": "⏳ AI: processing...", "ai.professional": "Professional tone", + "ai.provider": "Provider", + "ai.provider_default": "— default —", + "ai.model": "Model", + "ai.model_default": "— default —", + "ai.model_loading": "Loading…", + "ai.model_offline": "offline", + "ai.model_load_error": "Load error", "ai.quota_exceeded": "AI: quota exceeded or payment required", "ai.rewrite": "💬 Rewrite", "ai.rewrite_done": "AI: text rewritten", diff --git a/frontend/locales/fr.json b/frontend/locales/fr.json index 7c041ec..d3371f5 100644 --- a/frontend/locales/fr.json +++ b/frontend/locales/fr.json @@ -102,6 +102,13 @@ "ai.not_configured": "⚠️ AI non configuré — ajouter DEEPSEEK_API_KEY, OPENROUTER_API_KEY ou GEMINI_API_KEY dans .env", "ai.processing": "⏳ AI: traitement en cours...", "ai.professional": "Ton professionnel", + "ai.provider": "Fournisseur", + "ai.provider_default": "— défaut —", + "ai.model": "Modèle", + "ai.model_default": "— défaut —", + "ai.model_loading": "Chargement…", + "ai.model_offline": "hors ligne", + "ai.model_load_error": "Erreur de chargement", "ai.quota_exceeded": "AI: quota dépassé ou paiement requis", "ai.rewrite": "💬 Réécrire", "ai.rewrite_done": "AI: texte réécrit", diff --git a/tests/frontend/validate-imports.mjs b/tests/frontend/validate-imports.mjs index d09375a..d638690 100644 --- a/tests/frontend/validate-imports.mjs +++ b/tests/frontend/validate-imports.mjs @@ -38,8 +38,15 @@ function collectExports(filePath, modName) { const m = exportBlockText.match(/^export\s*\{([^}]+)\}/); if (m) { for (const name of m[1].split(',')) { - const n = name.trim().replace(/\s+as\s+\w+.*/, '').trim(); - if (n && n !== '') exports.add(n); + // Handle `original as exported` — export both names + const asMatch = name.trim().match(/^(\w+)\s+as\s+(\w+)$/); + if (asMatch) { + exports.add(asMatch[2]); // exported name + exports.add(asMatch[1]); // original name too (in case it's re-imported) + } else { + const n = name.trim().trim(); + if (n && n !== '') exports.add(n); + } } } exportBlockText = ''; diff --git a/tests/test_ai_models.py b/tests/test_ai_models.py new file mode 100644 index 0000000..5445b69 --- /dev/null +++ b/tests/test_ai_models.py @@ -0,0 +1,204 @@ +"""Tests for the AI models listing endpoint + curated fallback lists (ROADMAP #74/#71). + +Covers: +- GET /api/config/ai-models always returns a non-empty list for known providers, + even when the live API call fails (network, auth, etc.). +- The fallback list is curated with at least one model per provider. +- Gemini parsing strips the "models/" prefix. +- The default model is prepended if not already in the list. +""" +from __future__ import annotations + +import pytest +from fastapi.testclient import TestClient + +from backend.main import _FALLBACK_MODELS, app + + +@pytest.fixture +def admin_client(tmp_path): + """Minimal admin client for the /api/config/ai-models endpoint.""" + from backend.auth.password import hash_password + import json + import os + from pathlib import Path + + data_dir = tmp_path / "data" + data_dir.mkdir() + users = { + "version": 1, + "users": { + "admin": { + "id": "admin-1", + "username": "admin", + "display_name": "admin", + "password_hash": hash_password("chab30"), + "role": "admin", + "vaults": ["*"], + "active": True, + "created_at": "2026-01-01T00:00:00", + }, + }, + } + (data_dir / "users.json").write_text(json.dumps(users), encoding="utf-8") + + src_secret = Path("data/secret.key") + if src_secret.exists(): + import shutil + shutil.copy2(str(src_secret), str(data_dir / "secret.key")) + + orig_cwd = os.getcwd() + os.chdir(str(tmp_path)) + + os.environ["VAULT_1_NAME"] = "TestVault" + os.environ["VAULT_1_PATH"] = str(Path("test-vault").resolve()) + os.environ["OBSIGATE_AUTH_ENABLED"] = "true" + os.environ["OBSIGATE_ADMIN_USER"] = "admin" + os.environ["OBSIGATE_ADMIN_PASSWORD"] = "chab30" + os.environ["OBSIGATE_WATCHER_ENABLED"] = "false" + + import backend.main + backend.main._load_config = lambda: {"watcher_enabled": False} + from backend.indexer import build_index, index + import asyncio + for key in list(index.keys()): + del index[key] + loop = asyncio.new_event_loop() + asyncio.set_event_loop(loop) + loop.run_until_complete(build_index()) + + client = TestClient(app) + yield client + client.close() + os.chdir(orig_cwd) + for k in ["VAULT_1_NAME", "VAULT_1_PATH", "OBSIGATE_AUTH_ENABLED", + "OBSIGATE_ADMIN_USER", "OBSIGATE_ADMIN_PASSWORD", "OBSIGATE_WATCHER_ENABLED"]: + os.environ.pop(k, None) + + +def _login_admin(client): + resp = client.post("/api/auth/login", + json={"username": "admin", "password": "chab30"}) + assert resp.status_code == 200, resp.text + return resp.json()["access_token"] + + +# ── Unit tests for the fallback table itself ───────────────────────────── + + +class TestFallbackTable: + def test_every_known_provider_has_fallback_list(self): + """Every provider in the dropdown must have a non-empty fallback list.""" + expected_providers = { + "deepseek", "openrouter", "gemini", "nvidia", + "qwencloud", "xiaomi", "mistral", + } + assert set(_FALLBACK_MODELS.keys()) >= expected_providers + for p in expected_providers: + assert _FALLBACK_MODELS[p], f"fallback list for {p!r} is empty" + assert all(isinstance(m, str) and m.strip() for m in _FALLBACK_MODELS[p]), \ + f"fallback list for {p!r} contains invalid entries: {_FALLBACK_MODELS[p]}" + + def test_fallback_lists_are_short_and_focused(self): + """Fallbacks should be short (≤10 models) and well-known.""" + for p, models in _FALLBACK_MODELS.items(): + assert len(models) <= 10, f"too many fallbacks for {p}: {len(models)}" + + def test_xiaomi_fallback_contains_mimo_models(self): + """Xiaomi's MiMo models must be in the fallback list.""" + mimos = [m for m in _FALLBACK_MODELS["xiaomi"] if "mimo" in m.lower()] + assert mimos, "no MiMo model in xiaomi fallback" + + +# ── Endpoint integration tests ──────────────────────────────────────────── + + +class TestListModelsEndpoint: + """Verify /api/config/ai-models?provider=X always returns a usable list.""" + + def test_unknown_provider_returns_empty(self, admin_client): + token = _login_admin(admin_client) + resp = admin_client.get( + "/api/config/ai-models", + params={"provider": "nonexistent-provider"}, + headers={"Authorization": f"Bearer {token}"}, + ) + assert resp.status_code == 200 + data = resp.json() + assert data["models"] == [] + assert "error" in data or "source" in data + + def test_provider_without_key_returns_fallback(self, admin_client, monkeypatch): + """When the provider has no API key configured, return the curated fallback list.""" + # Force every provider key to be empty so the endpoint hits its + # 'no key configured' branch. + from backend import ai + monkeypatch.setattr(ai, "get_ai_key", lambda name: None) + + token = _login_admin(admin_client) + for provider in ("xiaomi", "nvidia", "deepseek", "mistral"): + resp = admin_client.get( + "/api/config/ai-models", + params={"provider": provider}, + headers={"Authorization": f"Bearer {token}"}, + ) + assert resp.status_code == 200, f"{provider}: {resp.status_code}" + data = resp.json() + assert data["models"], f"{provider}: fallback list is empty" + assert data.get("source") == "fallback", ( + f"{provider}: expected source=fallback, got {data.get('source')!r}" + ) + + def test_provider_with_unreachable_api_returns_fallback( + self, admin_client, monkeypatch, + ): + """When the live API call fails (network/DNS), the fallback list is used. + + This test patches both the key lookup AND the urlopen call so we + deterministically hit the network-failure branch. + """ + from backend import ai as aimod + + # Fake key so we don't take the 'no key' short-circuit. + monkeypatch.setattr(aimod, "get_ai_key", lambda name: "fake-key-for-test") + + # Patch urllib.request.urlopen to always raise — simulating network down. + # NOTE: main.py imports urllib.request at module load, so we patch the + # symbol it actually uses (urllib.request.urlopen). + import urllib.request + def _boom(*args, **kwargs): + raise OSError("simulated network down") + monkeypatch.setattr(urllib.request, "urlopen", _boom) + + token = _login_admin(admin_client) + # Pick a non-Gemini provider so we hit the network code path. + resp = admin_client.get( + "/api/config/ai-models", + params={"provider": "nvidia"}, + headers={"Authorization": f"Bearer {token}"}, + ) + assert resp.status_code == 200 + data = resp.json() + # The response should be a usable list — either via fallback after + # network failure, or the 'no key' shortcut. Both are acceptable as + # long as models is non-empty. + assert data["models"], "model list should be non-empty" + assert data.get("source") in ("fallback",), ( + f"unexpected source: {data.get('source')!r}" + ) + + def test_response_includes_source_field(self, admin_client, monkeypatch): + """All successful responses must include a 'source' field for UI hinting.""" + from backend import ai as aimod + monkeypatch.setattr(aimod, "get_ai_key", lambda name: "fake-key") + + token = _login_admin(admin_client) + resp = admin_client.get( + "/api/config/ai-models", + params={"provider": "gemini"}, + headers={"Authorization": f"Bearer {token}"}, + ) + assert resp.status_code == 200 + data = resp.json() + assert "source" in data + assert data["source"] in ("live", "fallback") diff --git a/tests/test_bookslm.py b/tests/test_bookslm.py index 9f82acb..c4a20a8 100644 --- a/tests/test_bookslm.py +++ b/tests/test_bookslm.py @@ -536,3 +536,41 @@ class TestBooksLMChatEndpoint: headers={"Authorization": f"Bearer {token}"}, ) assert resp.status_code == 404 + + def test_chat_request_accepts_provider_and_model(self, bookslm_client): + """The chat endpoint schema accepts provider + model fields without rejecting. + + We don't actually call the AI (would need a live key) — we just verify + the request schema is wired correctly so a missing key is the only + failure mode, not a 422 validation error. + """ + token, _ = _login_bookslm(bookslm_client) + # Build a tiny valid directory so the chat endpoint doesn't 404. + from pathlib import Path + vault_dir = Path(os.environ.get("VAULT_1_PATH", "test-vault")) + sub = vault_dir / "for_chat_test" + sub.mkdir(exist_ok=True) + (sub / "note.md").write_text("# hello\n", encoding="utf-8") + try: + resp = bookslm_client.post( + "/api/ai/bookslm/chat", + json={ + "vault": "TestVault", + "directory": "for_chat_test", + "message": "ping", + "provider": "deepseek", + "model": "deepseek-chat", + }, + headers={"Authorization": f"Bearer {token}"}, + ) + # Either 200 (if a real key is configured) or 500 (no key / quota). + # Must NOT be 422 — the schema must accept the fields. + assert resp.status_code in (200, 500), ( + f"unexpected status {resp.status_code}: {resp.text[:200]}" + ) + assert resp.status_code != 422, ( + f"schema rejected provider/model fields: {resp.text[:300]}" + ) + finally: + import shutil + shutil.rmtree(sub, ignore_errors=True)