Files
ObsiGate/backend/ai_routes.py
T
bruno a3642caa3d feat(ai): selection fournisseur/modele par defaut + guide architecture
- Persistance ai_default_provider / ai_default_models dans data/config.json
- Lecture + rechargement a chaud dans backend/ai.py (get_default_provider, reload_ai_config)
- UI: selecteurs Fournisseur/Modele par defaut dans la section Cles API IA (i18n FR/EN)
- docs/AI_ARCHITECTURE_GUIDE.md: architecture cible, catalogue d'outils, decisions MCP
2026-09-11 12:03:16 -04:00

235 lines
8.3 KiB
Python

"""ObsiGate AI — API routes for AI-powered editor features."""
import logging
from fastapi import APIRouter, HTTPException
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.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)