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.
175 lines
6.9 KiB
Python
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()
|
|
}
|