Files
ObsiGate/backend/model_capabilities.py
T
bruno f049e208b6
CI / lint (push) Successful in 1m10s
CI / security (push) Successful in 43s
CI / test (push) Successful in 2m34s
CI / build (push) Successful in 43s
CI / e2e (push) Successful in 10m59s
feat(ai): commandes @/ & skills, analyse d'images et capacites des modeles (#81)
2026-09-12 11:38:46 -04:00

145 lines
4.6 KiB
Python

"""Curated model-capability metadata for the AI assistant.
ObsiGate does not query every provider for the modalities a model supports
(not all of them expose that information, and the network call is not always
reliable). Instead a static, curated table maps known model-name patterns to
capability flags, with per-provider defaults. The UI uses this to show, when a
model is selected, which features it supports (Chat, Embeddings, Rerank,
Images, Video, Audio Speech, Audio Transcriptions, Vision).
The table is intentionally conservative: an unknown model falls back to the
provider default (usually ``chat`` only), so we never claim a capability the
model may not have.
"""
from __future__ import annotations
from typing import Any
# Ordered list of capability keys exposed to the UI. Keep in sync with the
# frontend ``AI_CAPABILITY_KEYS`` and the i18n ``ai.cap_*`` labels.
CAPABILITY_KEYS: tuple[str, ...] = (
"chat",
"embeddings",
"rerank",
"images",
"video",
"audio_speech",
"audio_transcription",
"vision",
)
def _caps(**kwargs: bool) -> dict[str, bool]:
"""Build a full capability dict (missing flags default to False)."""
base = {key: False for key in CAPABILITY_KEYS}
base.update(kwargs)
return base
# Provider defaults, used when no model-name pattern matches.
_PROVIDER_DEFAULTS: dict[str, dict[str, bool]] = {
"deepseek": _caps(chat=True),
"openrouter": _caps(chat=True),
"gemini": _caps(chat=True, vision=True, embeddings=True, audio_speech=True),
"ollama": _caps(chat=True, embeddings=True),
"nvidia": _caps(chat=True),
"qwencloud": _caps(chat=True),
"xiaomi": _caps(chat=True),
"mistral": _caps(chat=True, embeddings=True),
}
# Ordered (substrings, capabilities) rules — the first matching rule wins.
# More specific modalities are listed before the broad vision/chat rule.
_MODEL_RULES: list[tuple[tuple[str, ...], dict[str, bool]]] = [
# Rerankers.
(("rerank", "cross-encoder"), _caps(rerank=True)),
# Embedding models (usually not chat-capable).
(
("text-embedding", "embed", "bge-", "e5-", "nomic-embed", "gte-"),
_caps(embeddings=True),
),
# Speech-to-text / transcription.
(
("whisper", "transcrib", "-asr", "asr-", "speech-to-text"),
_caps(audio_transcription=True),
),
# Text-to-speech.
(
("-tts", "text-to-speech", "voiceclone", "voicedesign"),
_caps(audio_speech=True),
),
# Image generation.
(
("dall-e", "stable-diffusion", "flux", "imagen", "image-gen"),
_caps(images=True),
),
# Video generation.
(("veo-", "sora", "video-gen", "-video"), _caps(video=True)),
# Vision-capable chat models (multimodal input).
(
(
"gpt-4o",
"gpt-4.1",
"gpt-5",
"o3",
"o4",
"claude-3",
"claude-4",
"gemini",
"qwen-vl",
"-vl-",
"vl-",
"vision",
"llava",
"pixtral",
"internvl",
"minicpm-v",
"llama-3.2-vision",
"mimo-vl",
),
_caps(chat=True, vision=True),
),
]
def get_model_capabilities(provider: str, model: str) -> dict[str, bool]:
"""Return the capability flags for a ``provider``/``model`` pair.
Args:
provider: Provider identifier (e.g. ``"deepseek"``). Case-insensitive.
model: Model identifier (e.g. ``"deepseek-chat"``). May be empty, in
which case the provider default is returned.
Returns:
A dict with every key of :data:`CAPABILITY_KEYS` and boolean values.
"""
provider = (provider or "").strip().lower()
name = (model or "").strip().lower()
if name:
for needles, caps in _MODEL_RULES:
if any(needle in name for needle in needles):
return dict(caps)
return dict(_PROVIDER_DEFAULTS.get(provider, _caps(chat=True)))
def get_capabilities_for_models(
provider: str, models: list[str]
) -> dict[str, dict[str, bool]]:
"""Return a ``{model: capabilities}`` map for a list of models."""
return {model: get_model_capabilities(provider, model) for model in models}
def model_supports_vision(provider: str, model: str) -> bool:
"""True when the given model can accept image input."""
return bool(get_model_capabilities(provider, model).get("vision"))
def capabilities_payload(provider: str, model: str) -> dict[str, Any]:
"""Serialize capabilities for API responses."""
return {
"provider": provider,
"model": model,
"capabilities": get_model_capabilities(provider, model),
}