Le panneau de modèle par défaut et la bulle ⓘ n'affichaient aucun modèle Mistral « Vision capable » alors que GET api.mistral.ai/v1/models en déclare 28 : la table de capacités était entièrement statique et aucun de ses motifs ne correspondait aux familles Mistral actuelles (seul `pixtral`, retiré de l'API, les matchait). - backend/provider_capabilities.py (nouveau) : capacités déclarées par le fournisseur (Mistral `capabilities`, OpenRouter `architecture`), détectées par la forme du payload, snapshot en cache process-wide (TTL 30 min, surchargeable par AI_CAPABILITIES_TTL_SECONDS) rempli par GET /api/config/ai-models. - backend/model_capabilities.py : une déclaration prime sur la table statique pour chaque drapeau mentionné ; la table ne comble que le reste (Mistral ne déclare jamais `embedding`). Table corrigée pour le repli hors ligne : familles vision Mistral (ministral, magistral, mistral-small, mistral-medium, mistral-vibe-cli, labs-leanstral), mistral-ocr = vision sans chat, et défaut du fournisseur Mistral sans `embeddings` (mistral-large / codestral n'étaient plus des « embedders »). - backend/ai_routes.py : GET /api/ai/model-capabilities reste sans appel réseau (cache froid → table statique). - Tests : tests/test_provider_capabilities.py (nouveau), TestMistralFamilies et TestDeclaredCapabilities (bout en bout via l'API). - Docs : CHANGELOG [Unreleased], registre + journal ISSUES_TODOLIST, fiche docs/features/ai-provider-picker.md (§L). Vérifié : 28/28 modèles vision déclarés par Mistral détectés (0 avant), 0 écart dans les deux sens ; pytest 1007 passed / 6 skipped ; ruff 0 (backend) ; mypy 0 ; tests frontend unit 9/9 + IA 66/66 + validate-imports 37 modules ; instance de test reconstruite et vérifiée sur http://localhost:2020.
270 lines
9.6 KiB
Python
270 lines
9.6 KiB
Python
"""ObsiGate AI — API routes for AI-powered editor features."""
|
|
|
|
import logging
|
|
|
|
from fastapi import APIRouter, Depends, HTTPException, Query
|
|
from pydantic import BaseModel, Field
|
|
|
|
from backend.ai import (
|
|
PROVIDERS,
|
|
ai_change_tone,
|
|
ai_continue_writing,
|
|
ai_convert_to_canvas,
|
|
ai_convert_to_list,
|
|
ai_convert_to_table,
|
|
ai_custom_rewrite,
|
|
ai_explain,
|
|
ai_fix_spelling,
|
|
ai_generate_frontmatter,
|
|
ai_improve_writing,
|
|
ai_inline_complete,
|
|
ai_make_longer,
|
|
ai_make_shorter,
|
|
ai_simplify,
|
|
ai_summarize,
|
|
ai_translate,
|
|
get_default_provider,
|
|
)
|
|
from backend.auth.middleware import require_auth
|
|
from backend.model_capabilities import get_model_capabilities
|
|
from backend.schemas import AIStatusResponse
|
|
|
|
logger = logging.getLogger("obsigate.ai_routes")
|
|
router = APIRouter(prefix="/api/ai", tags=["AI"])
|
|
|
|
|
|
@router.get("/status", response_model=AIStatusResponse)
|
|
async def api_status():
|
|
"""Check if AI is configured and which providers are available.
|
|
|
|
Reads environment variables at runtime so that changes to
|
|
.env or Docker environment are reflected immediately.
|
|
"""
|
|
from backend.ai import get_ai_key
|
|
provider_keys = {
|
|
"deepseek": "DEEPSEEK_API_KEY",
|
|
"openrouter": "OPENROUTER_API_KEY",
|
|
"gemini": "GEMINI_API_KEY",
|
|
"nvidia": "NVIDIA_API_KEY",
|
|
"qwencloud": "QWENCLOUD_API_KEY",
|
|
"xiaomi": "XIAOMI_API_KEY",
|
|
"mistral": "MISTRAL_API_KEY",
|
|
}
|
|
providers = {}
|
|
for name, env_var in provider_keys.items():
|
|
has_key = bool(get_ai_key(env_var))
|
|
providers[name] = {
|
|
"available": has_key,
|
|
"model": PROVIDERS[name]["model"] if has_key else None,
|
|
}
|
|
|
|
# ---- Autocomplete (Ollama) status ----
|
|
ollama_cfg = PROVIDERS.get("ollama", {})
|
|
ollama_url = (ollama_cfg.get("base_url") or "").rstrip("/v1").rstrip("/")
|
|
ollama_model = ollama_cfg.get("model", "")
|
|
autocomplete = {
|
|
"available": False,
|
|
"server_ok": False,
|
|
"model_loaded": False,
|
|
"model": ollama_model,
|
|
"error": None,
|
|
}
|
|
if ollama_url:
|
|
try:
|
|
import httpx
|
|
async with httpx.AsyncClient(timeout=5.0) as client:
|
|
# Check server health + list loaded models
|
|
r = await client.get(f"{ollama_url}/api/tags")
|
|
if r.status_code == 200:
|
|
autocomplete["server_ok"] = True
|
|
data = r.json()
|
|
models = [m.get("name", "") for m in data.get("models", [])]
|
|
# Check if our model (with or without tag) is loaded
|
|
model_base = ollama_model.split(":")[0]
|
|
autocomplete["model_loaded"] = any(
|
|
m == ollama_model or m.startswith(model_base + ":")
|
|
for m in models
|
|
)
|
|
autocomplete["available"] = autocomplete["model_loaded"]
|
|
else:
|
|
autocomplete["error"] = f"HTTP {r.status_code}"
|
|
except Exception as e:
|
|
autocomplete["error"] = str(e)
|
|
|
|
return {
|
|
"configured": any(p["available"] for p in providers.values()),
|
|
"default_provider": get_default_provider(),
|
|
"providers": providers,
|
|
"autocomplete": autocomplete,
|
|
}
|
|
|
|
|
|
class AIRequest(BaseModel):
|
|
text: str = Field(..., description="Input text to process", min_length=1)
|
|
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 (e.g. 'deepseek', 'nvidia')")
|
|
model: str | None = Field(None, description="Model name override for this request")
|
|
|
|
|
|
class AIResponse(BaseModel):
|
|
result: str = Field(..., description="Processed text result")
|
|
provider: str = Field(..., description="AI provider used")
|
|
|
|
|
|
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")
|
|
except ValueError as e:
|
|
raise HTTPException(status_code=400, detail=str(e))
|
|
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)
|
|
async def api_improve(req: AIRequest):
|
|
"""Improve writing quality."""
|
|
return await _handle(ai_improve_writing, req)
|
|
|
|
|
|
@router.post("/fix-spelling", response_model=AIResponse)
|
|
async def api_fix_spelling(req: AIRequest):
|
|
"""Fix spelling and grammar."""
|
|
return await _handle(ai_fix_spelling, req)
|
|
|
|
|
|
@router.post("/make-shorter", response_model=AIResponse)
|
|
async def api_make_shorter(req: AIRequest):
|
|
"""Make text more concise."""
|
|
return await _handle(ai_make_shorter, req)
|
|
|
|
|
|
@router.post("/make-longer", response_model=AIResponse)
|
|
async def api_make_longer(req: AIRequest):
|
|
"""Expand text with more detail."""
|
|
return await _handle(ai_make_longer, req)
|
|
|
|
|
|
@router.post("/simplify", response_model=AIResponse)
|
|
async def api_simplify(req: AIRequest):
|
|
"""Simplify language."""
|
|
return await _handle(ai_simplify, req)
|
|
|
|
|
|
@router.post("/tone", response_model=AIResponse)
|
|
async def api_tone(req: AIRequest):
|
|
"""Change text tone. Requires `tone` field (e.g., 'professional', 'casual')."""
|
|
if not req.tone:
|
|
raise HTTPException(status_code=400, detail="Field 'tone' is required (e.g., 'professional', 'casual')")
|
|
return await _handle(lambda text, p: ai_change_tone(text, req.tone, p), req)
|
|
|
|
|
|
@router.post("/translate", response_model=AIResponse)
|
|
async def api_translate(req: AIRequest):
|
|
"""Translate text. Requires `target_lang` field (e.g., 'French', 'English', 'Japanese')."""
|
|
if not req.target_lang:
|
|
raise HTTPException(status_code=400, detail="Field 'target_lang' is required")
|
|
return await _handle(lambda text, p: ai_translate(text, req.target_lang, p), req)
|
|
|
|
|
|
@router.post("/explain", response_model=AIResponse)
|
|
async def api_explain(req: AIRequest):
|
|
"""Explain the selected text."""
|
|
return await _handle(ai_explain, req)
|
|
|
|
|
|
@router.post("/summarize", response_model=AIResponse)
|
|
async def api_summarize(req: AIRequest):
|
|
"""Summarize text."""
|
|
return await _handle(ai_summarize, req)
|
|
|
|
|
|
@router.post("/continue", response_model=AIResponse)
|
|
async def api_continue(req: AIRequest):
|
|
"""Continue writing from the selected text."""
|
|
return await _handle(ai_continue_writing, req)
|
|
|
|
|
|
@router.post("/rewrite", response_model=AIResponse)
|
|
async def api_rewrite(req: AIRequest):
|
|
"""Custom rewrite with instruction. Requires `instruction` field."""
|
|
if not req.instruction:
|
|
raise HTTPException(status_code=400, detail="Field 'instruction' is required")
|
|
return await _handle(lambda text, p: ai_custom_rewrite(text, req.instruction, p), req)
|
|
|
|
|
|
@router.post("/to-list", response_model=AIResponse)
|
|
async def api_to_list(req: AIRequest):
|
|
"""Convert text to a markdown list."""
|
|
return await _handle(ai_convert_to_list, req)
|
|
|
|
|
|
@router.post("/to-table", response_model=AIResponse)
|
|
async def api_to_table(req: AIRequest):
|
|
"""Convert text to a markdown table."""
|
|
return await _handle(ai_convert_to_table, req)
|
|
|
|
|
|
@router.post("/frontmatter", response_model=AIResponse)
|
|
async def api_frontmatter(req: AIRequest):
|
|
"""Generate YAML frontmatter."""
|
|
return await _handle(ai_generate_frontmatter, req)
|
|
|
|
|
|
@router.post("/inline-complete", response_model=AIResponse)
|
|
async def api_inline_complete(req: AIRequest):
|
|
"""Inline completion."""
|
|
return await _handle(ai_inline_complete, req)
|
|
|
|
|
|
@router.post("/to-canvas", response_model=AIResponse)
|
|
async def api_to_canvas(req: AIRequest):
|
|
"""Convert to Mermaid diagram or outline."""
|
|
return await _handle(ai_convert_to_canvas, req)
|
|
|
|
|
|
class ModelCapabilitiesResponse(BaseModel):
|
|
"""Capabilities of a single provider/model pair."""
|
|
|
|
provider: str
|
|
model: str
|
|
capabilities: dict[str, bool] = Field(
|
|
description="Flags: chat, embeddings, rerank, images, video, "
|
|
"audio_speech, audio_transcription, vision",
|
|
)
|
|
|
|
|
|
@router.get("/model-capabilities", response_model=ModelCapabilitiesResponse)
|
|
async def api_model_capabilities(
|
|
provider: str = Query(..., description="Provider identifier"),
|
|
model: str = Query("", description="Model identifier (optional)"),
|
|
current_user=Depends(require_auth),
|
|
):
|
|
"""Return the capability flags for a provider/model pair.
|
|
|
|
Two layers (BUG-044): flags the provider itself declares in its models
|
|
endpoint (Mistral ``capabilities``, OpenRouter ``architecture``) win, the
|
|
curated table in ``backend.model_capabilities`` fills the rest. The
|
|
declaration snapshot is populated by ``GET /api/config/ai-models``; when it
|
|
is cold (or the provider declares nothing) the curated table answers alone,
|
|
so this endpoint never performs a blocking provider call.
|
|
"""
|
|
return {
|
|
"provider": provider,
|
|
"model": model,
|
|
"capabilities": get_model_capabilities(provider, model),
|
|
}
|