145 lines
4.6 KiB
Python
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),
|
|
}
|