"""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() }