Files
ObsiGate/backend/provider_capabilities.py
T
bruno b69cb9b0f8
CI / lint (push) Successful in 1m20s
CI / security (push) Successful in 53s
CI / test (push) Successful in 2m27s
CI / build (push) Successful in 50s
CI / e2e (push) Successful in 10m48s
fix(ai): BUG-044 capacités des modèles lues chez le fournisseur (Mistral vision)
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.
2026-09-15 09:26:31 -04:00

175 lines
6.9 KiB
Python

"""Provider-declared model capabilities (live) with a short-lived cache.
The curated table in :mod:`backend.model_capabilities` has to be edited by hand
every time a provider ships or renames a model, and it ages badly: Mistral alone
declares ``vision`` on 28 of its models while the curated table knew none of
them (BUG-044). Providers that *do* publish per-model capabilities in their
public models endpoint are therefore asked first; the curated table is only
used to fill the flags the provider stays silent about (e.g. Mistral never
declares ``embedding``, only the absence of ``completion_chat``).
Supported declarations, detected by *payload shape* so a provider that starts
exposing them is picked up without a code change:
* ``capabilities`` dict (Mistral): ``completion_chat`` → ``chat``, ``vision``,
``audio_transcription`` (+ ``audio_transcription_realtime``),
``audio_speech``.
* ``architecture`` dict (OpenRouter): ``input_modalities`` /
``output_modalities`` → ``vision`` (image input), ``images`` (image output),
``audio_transcription`` (audio input), ``audio_speech`` (audio output),
``video`` (video output), ``chat`` (text output).
The cache is in-process and shared by every request. Its TTL only bounds how
long a *stale* declaration can survive: a fresh provider call (the model list
endpoint) overwrites the provider entry immediately.
"""
from __future__ import annotations
import logging
import os
import time
from typing import Any
logger = logging.getLogger("obsigate.ai.capabilities")
#: How long a provider-declared snapshot stays usable, in seconds.
#: ``0`` disables expiry (the snapshot lives until the next provider call).
TTL_SECONDS = float(os.getenv("AI_CAPABILITIES_TTL_SECONDS", "1800") or 0)
#: Guard against a provider returning a runaway model list.
MAX_MODELS_PER_PROVIDER = 2000
# provider → (timestamp, model count, {normalized model id: declared flags})
_CACHE: dict[str, tuple[float, int, dict[str, dict[str, bool]]]] = {}
def _norm(value: str) -> str:
"""Lower-case and strip the ``models/`` prefix Gemini uses."""
return (value or "").strip().lower().removeprefix("models/")
def _from_capability_flags(flags: dict[str, Any]) -> dict[str, bool]:
"""Map a Mistral-style ``capabilities`` dict onto ObsiGate flags."""
declared: dict[str, bool] = {}
if isinstance(flags.get("completion_chat"), bool):
declared["chat"] = flags["completion_chat"]
if isinstance(flags.get("vision"), bool):
declared["vision"] = flags["vision"]
transcription = flags.get("audio_transcription")
realtime = flags.get("audio_transcription_realtime")
if isinstance(transcription, bool) or isinstance(realtime, bool):
declared["audio_transcription"] = bool(transcription or realtime)
if isinstance(flags.get("audio_speech"), bool):
declared["audio_speech"] = flags["audio_speech"]
return declared
def _from_architecture(architecture: dict[str, Any]) -> dict[str, bool]:
"""Map an OpenRouter-style ``architecture`` dict onto ObsiGate flags."""
inputs = architecture.get("input_modalities")
outputs = architecture.get("output_modalities")
if not isinstance(inputs, list) and not isinstance(outputs, list):
return {}
in_modalities = [str(m).lower() for m in inputs] if isinstance(inputs, list) else []
out_modalities = [str(m).lower() for m in outputs] if isinstance(outputs, list) else []
return {
"chat": "text" in out_modalities,
"vision": "image" in in_modalities,
"images": "image" in out_modalities,
"audio_transcription": "audio" in in_modalities,
"audio_speech": "audio" in out_modalities,
"video": "video" in out_modalities,
}
def parse_declared_capabilities(entry: Any) -> dict[str, bool] | None:
"""Extract the capability flags a single provider model entry declares.
Args:
entry: One item of a provider models payload (``/v1/models``).
Returns:
A partial ``{flag: bool}`` mapping (only the flags the provider
actually declares), or ``None`` when the entry declares nothing.
"""
if not isinstance(entry, dict):
return None
flags = entry.get("capabilities")
declared = _from_capability_flags(flags) if isinstance(flags, dict) else {}
if not declared:
architecture = entry.get("architecture")
declared = _from_architecture(architecture) if isinstance(architecture, dict) else {}
return declared or None
def remember_declared_capabilities(provider: str, payload: Any) -> int:
"""Cache the capabilities declared by a provider models payload.
Args:
provider: Provider identifier (e.g. ``"mistral"``).
payload: Raw JSON body of the provider models endpoint, or the model
list itself.
Returns:
The number of models with declared capabilities that were cached.
"""
provider = (provider or "").strip().lower()
entries: Any = payload.get("data") if isinstance(payload, dict) else payload
if not isinstance(entries, list):
return 0
parsed: dict[str, dict[str, bool]] = {}
for entry in entries[:MAX_MODELS_PER_PROVIDER]:
if not isinstance(entry, dict):
continue
model_id = entry.get("id") or entry.get("name") or ""
if not isinstance(model_id, str) or not model_id.strip():
continue
declared = parse_declared_capabilities(entry)
if declared:
parsed[_norm(model_id)] = declared
if not parsed:
return 0
_CACHE[provider] = (time.time(), len(parsed), parsed)
logger.info(f"Capabilities declared by {provider}: {len(parsed)} models cached")
return len(parsed)
def get_declared_capabilities(provider: str, model: str) -> dict[str, bool] | None:
"""Return the cached declared capabilities for one provider/model pair.
Returns ``None`` when nothing was declared for that pair (cache cold,
expired, or the provider is silent about this model).
"""
snapshot = _CACHE.get((provider or "").strip().lower())
if not snapshot:
return None
timestamp, _count, table = snapshot
if TTL_SECONDS and (time.time() - timestamp) > TTL_SECONDS:
return None
declared = table.get(_norm(model))
return dict(declared) if declared else None
def clear_declared_capabilities(provider: str | None = None) -> None:
"""Drop the cached snapshot for one provider, or all of them (tests/ops)."""
if provider is None:
_CACHE.clear()
return
_CACHE.pop(provider.strip().lower(), None)
def cache_info() -> dict[str, dict[str, Any]]:
"""Diagnostics: per-provider cache age and model count."""
now = time.time()
return {
provider: {
"models": count,
"age_seconds": round(now - timestamp, 1),
"expired": bool(TTL_SECONDS and (now - timestamp) > TTL_SECONDS),
}
for provider, (timestamp, count, _table) in _CACHE.items()
}