diff --git a/agent/models_dev.py b/agent/models_dev.py index 7d197ec1db..be418d337e 100644 --- a/agent/models_dev.py +++ b/agent/models_dev.py @@ -1,27 +1,17 @@ """Models.dev registry integration — primary database for providers and models. -Fetches https://models.dev/api.json (4000+ models, 100+ providers): provider -metadata (name, base URL, env vars, docs) and model metadata (context window, -max output, cost/M tokens, capabilities, modalities, knowledge cutoff, -open-weights flag, family, deprecation status). - -Resolution order: in-memory cache (fresh, or stale served immediately while -one background daemon thread refreshes) → disk cache -(~/.hermes/models_dev_cache.json, any age) → network, only when no cache -exists at all. Failed refreshes back off for 5 minutes process-wide. - -Invariants: -- **ETag conditional GET**: refreshes send ``If-None-Match`` whenever a - servable registry is held; a 304 re-confirms the cache without - re-downloading ~2 MB. The ETag sidecar is persisted with the cache body. -- **No network on hot paths**: resolution, picker, and resume paths pass - ``allow_network=False`` and never perform network I/O. -- **Corrupt-cache rejection**: a disk cache that fails to parse, is not a - dict, or is empty is quarantined with a warning rather than served as ``{}``. -- **Mirror URL override**: ``models_dev.url`` in config.yaml. - -Other modules should import the dataclasses and query functions from here -rather than parsing the raw JSON themselves. +Fetches https://models.dev/api.json (provider + model metadata). Resolution +order: in-memory cache (fresh, or stale served immediately while one background +daemon thread refreshes) → disk cache (~/.hermes/models_dev_cache.json, any +age) → network, only when no cache exists at all. Failed refreshes back off for +5 minutes process-wide. Refreshes use ETag conditional GET whenever a servable +registry is held (a 304 re-confirms the cache without re-downloading ~2 MB); +the sidecar is persisted with the cache body. Hot paths (resolution, picker, +resume) pass ``allow_network=False`` and never do network I/O. A disk cache +that fails to parse, is not a dict, or is empty is quarantined with a warning +rather than served as ``{}``. ``models_dev.url`` in config.yaml overrides the +URL (mirrors). Other modules should use the dataclasses and query functions +here rather than parsing the raw JSON. """ import json @@ -33,7 +23,7 @@ from dataclasses import dataclass from pathlib import Path from typing import Any, Dict, List, Optional, Tuple -from utils import atomic_json_write +from utils import atomic_json_write, atomic_write_text import requests @@ -209,18 +199,34 @@ def _models_dev_to_hermes_ids(mdev_id: str) -> List[str]: return _MODELS_DEV_TO_PROVIDER.get(mdev_id, []) +def _dict_or_empty(value: Any) -> Dict[str, Any]: + return value if isinstance(value, dict) else {} + + +def _cfg_get(*keys: str, default: Any) -> Any: + """``cfg_get`` over the read-only config; *default* on any failure.""" + try: + from hermes_cli.config import cfg_get, load_config_readonly + return cfg_get(load_config_readonly(), *keys, default=default) + except Exception: + return default + + # --------------------------------------------------------------------------- # Disk cache + ETag sidecar # --------------------------------------------------------------------------- -def _get_cache_path() -> Path: +def _hermes_path(name: str) -> Path: from hermes_constants import get_hermes_home - return get_hermes_home() / "models_dev_cache.json" + return get_hermes_home() / name + + +def _get_cache_path() -> Path: + return _hermes_path("models_dev_cache.json") def _get_etag_path() -> Path: - from hermes_constants import get_hermes_home - return get_hermes_home() / "models_dev_cache.etag" + return _hermes_path("models_dev_cache.etag") def _load_etag() -> str: @@ -236,8 +242,6 @@ def _load_etag() -> str: def _save_etag(etag: str) -> None: try: - from utils import atomic_write_text - etag_path = _get_etag_path() etag_path.parent.mkdir(parents=True, exist_ok=True) atomic_write_text(etag_path, etag) @@ -248,8 +252,7 @@ def _save_etag(etag: str) -> None: def _clear_etag() -> None: """Delete the ETag sidecar so the next fetch is unconditional. - Called when the registry the ETag vouches for is gone or unusable: an - If-None-Match without a servable cache invites a 304 that leaves the + An If-None-Match without a servable cache invites a 304 that leaves the process with no data at all. """ try: @@ -259,18 +262,11 @@ def _clear_etag() -> None: def _get_models_dev_url() -> str: - """The models.dev API URL, honoring the ``models_dev.url`` config override - (mirrors / self-hosted copies).""" - try: - from hermes_cli.config import cfg_get, load_config_readonly - cfg = load_config_readonly() - url = cfg_get(cfg, "models_dev", "url", default="") - if isinstance(url, str) and url.strip(): - return url.strip() - except Exception: - pass - # Module global (not a captured constant) so code/tests that patch - # MODELS_DEV_URL keep working. + """The models.dev API URL, honoring the ``models_dev.url`` config override.""" + url = _cfg_get("models_dev", "url", default="") + if isinstance(url, str) and url.strip(): + return url.strip() + # Module global (not a captured constant) so patching MODELS_DEV_URL works. return MODELS_DEV_URL @@ -287,14 +283,13 @@ def _load_disk_cache() -> Dict[str, Any]: if cache_path.exists(): with open(cache_path, encoding="utf-8") as f: data = json.load(f) - if not _validate_registry(data): - logger.warning( - "models.dev disk cache is corrupt or empty; " - "quarantining (will refetch from network)" - ) - _quarantine_corrupt_cache(cache_path) - return {} - return data + if _validate_registry(data): + return data + logger.warning( + "models.dev disk cache is corrupt or empty; " + "quarantining (will refetch from network)" + ) + _quarantine_corrupt_cache(cache_path) except Exception as e: logger.warning( "Failed to load models.dev disk cache; quarantining: %s", e @@ -310,9 +305,9 @@ def _quarantine_corrupt_cache(cache_path: Path) -> None: """Rename a rejected cache aside and drop its ETag sidecar. Renaming makes the rejection a one-time event — otherwise every hot-path - call that finds the in-memory cache empty re-parses the corrupt file and - re-warns until a network fetch succeeds. The sidecar vouches for a registry - we no longer hold, so it goes too. + call that finds the in-memory cache empty re-parses and re-warns until a + network fetch succeeds. The sidecar vouches for a registry we no longer + hold, so it goes too. """ try: cache_path.rename(cache_path.with_suffix(".json.corrupt")) @@ -322,9 +317,11 @@ def _quarantine_corrupt_cache(cache_path: Path) -> None: def _disk_cache_age_seconds() -> Optional[float]: - """Age of the disk cache file in seconds, or None if missing/unreadable - (or mtime in the future from clock skew — treated as unknown freshness so - callers fall through to the network rather than trusting it forever).""" + """Age of the disk cache file in seconds, or None if missing/unreadable. + + An mtime in the future (clock skew) is also None — unknown freshness — so + callers fall through to the network rather than trusting it forever. + """ try: cache_path = _get_cache_path() if not cache_path.exists(): @@ -362,29 +359,24 @@ def _fetch_models_dev_from_network( ``conditional`` sends ``If-None-Match`` with the sidecar's ETag and raises ``_NotModified`` on 304. Pass True ONLY while holding ``_models_dev_fetch_lock`` AND a servable registry — a conditional request - without one invites a 304 that leaves the process with no data (formerly a - permanent empty-registry loop when the sidecar outlived a corrupt cache). + without one invites a 304 that leaves the process with no data. Raises on network errors and on an empty/invalid payload. """ - url = _get_models_dev_url() headers: Dict[str, str] = {} if conditional: etag = _load_etag() if etag: headers["If-None-Match"] = etag - # (connect, read) timeout: 5 s connect fails fast on blackholed hosts - # instead of stalling the first turn; 10 s read tolerates a slow registry. - response = requests.get(url, headers=headers, timeout=(5, 10)) - + # (connect, read): 5 s connect fails fast on blackholed hosts instead of + # stalling the first turn; 10 s read tolerates a slow registry. + response = requests.get(_get_models_dev_url(), headers=headers, timeout=(5, 10)) if response.status_code == 304: raise _NotModified() - response.raise_for_status() data = response.json() if not _validate_registry(data): raise ValueError("models.dev returned an empty or invalid registry") - return data, response.headers.get("ETag", "") @@ -401,6 +393,14 @@ def _mark_stale_cache_grace() -> None: _models_dev_cache_time = grace_time +def _serve_stale(msg: str, *args: Any) -> Dict[str, Any]: + """Arm the grace window, kick off a background refresh, return the held cache.""" + _mark_stale_cache_grace() + _start_background_refresh_models_dev() + logger.debug(msg, *args) + return _models_dev_cache + + def _commit_registry(data: Dict[str, Any], *, etag: str = "", where: str) -> None: """Persist a fetched registry: disk + in-mem + clear backoff. @@ -425,10 +425,9 @@ def _confirm_cache_not_modified(*, where: str) -> None: untouched — only the freshness marker advances). Caller holds the lock.""" global _models_dev_cache_time, _models_dev_retry_after if not _models_dev_cache: - # A 304 with no registry held should be unreachable (conditional GETs - # require a servable cache) but previously caused a permanent - # empty-registry loop: drop the sidecar and arm the normal backoff - # rather than marking {} "fresh". + # Should be unreachable (conditional GETs require a servable cache) but + # previously caused a permanent empty-registry loop: drop the sidecar + # and arm the normal backoff rather than marking {} "fresh". _clear_etag() _models_dev_retry_after = time.time() + _MODELS_DEV_RETRY_DELAY logger.warning( @@ -512,94 +511,69 @@ def fetch_models_dev( ) -> Dict[str, Any]: """Fetch the models.dev registry (dict keyed by provider ID; {} on failure). - Cache hierarchy when ``force_refresh=False``: - 1. Fresh in-memory cache → return. - 2. Stale in-memory cache → return it and refresh in one background daemon - thread. Callers never block on the network while any cache exists; - models.dev only changes when providers add models, so stale data beats - a foreground timeout. - 3. Disk cache (any age) → populate in-mem and return; a stale one - triggers the same background refresh. Corrupt/empty is rejected. - 4. No cache → singleflight foreground network fetch, saved to disk+mem. - 5. Any failed refresh suppresses automatic refreshes for 5 minutes. + Cache hierarchy when ``force_refresh=False``: fresh in-memory → stale + in-memory (returned immediately, refreshed in one background daemon thread + — stale data beats a foreground timeout) → disk cache of any age (a stale + one triggers the same background refresh; corrupt/empty is rejected) → + singleflight foreground network fetch. Any failed refresh suppresses + automatic refreshes for 5 minutes. ``force_refresh=True`` (``hermes config refresh``) bypasses the cache fast paths and the backoff, falling back to cached data only if the call fails. ``allow_network=False`` returns any memory/disk cache regardless of age and - never makes a request — for latency-sensitive paths (gateway route-identity - checks, vision routing, context-length lookup). - - Network requests use ETag conditional GET when a servable registry is held - (a cold ``force_refresh`` hydrates memory from disk first); a 304 - re-confirms the cache without re-downloading ~2 MB. + never makes a request — for latency-sensitive paths. """ global _models_dev_cache, _models_dev_cache_time, _models_dev_retry_after if not allow_network: - if _models_dev_cache: - return _models_dev_cache - disk_data = _load_disk_cache() - if disk_data: - _models_dev_cache = disk_data - disk_age = _disk_cache_age_seconds() - _models_dev_cache_time = ( - time.time() - disk_age if disk_age is not None else 0 - ) - return _models_dev_cache - - # Stage 1: fresh in-memory cache — the hot path, no I/O. - if ( - not force_refresh - and _models_dev_cache - and (time.time() - _models_dev_cache_time) < _MODELS_DEV_CACHE_TTL - ): - return _models_dev_cache - - # Stage 2: stale in-memory cache beats blocking on the network. - if not force_refresh and _models_dev_cache: - _mark_stale_cache_grace() - _start_background_refresh_models_dev() - logger.debug( - "Using stale in-memory models.dev cache; refreshing in background" - ) - return _models_dev_cache - - # Stage 3: disk cache (cold-start only). A stale disk cache is deliberately - # usable so resolution doesn't hang when models.dev is unreachable. - if not force_refresh: - disk_age = _disk_cache_age_seconds() - if disk_age is not None: + if not _models_dev_cache: disk_data = _load_disk_cache() if disk_data: _models_dev_cache = disk_data - if disk_age < _MODELS_DEV_CACHE_TTL: - # Anchor the in-mem TTL to the file's age so an aging - # cache isn't extended by another full TTL. - _models_dev_cache_time = time.time() - disk_age - logger.debug( - "Loaded models.dev from fresh disk cache " - "(%d providers, age=%.0fs)", len(disk_data), disk_age, - ) - else: - _mark_stale_cache_grace() - _start_background_refresh_models_dev() - logger.debug( - "Using stale models.dev disk cache (age=%.0fs); " - "refreshing in background", - disk_age, - ) - return _models_dev_cache - - # Process-wide backoff: don't make every caller retry an unreachable - # endpoint while no usable cache exists. - if not force_refresh and time.time() < _models_dev_retry_after: + disk_age = _disk_cache_age_seconds() + _models_dev_cache_time = ( + time.time() - disk_age if disk_age is not None else 0 + ) return _models_dev_cache + if not force_refresh: + # Stage 1: fresh in-memory cache — the hot path, no I/O. + if _models_dev_cache and (time.time() - _models_dev_cache_time) < _MODELS_DEV_CACHE_TTL: + return _models_dev_cache + # Stage 2: stale in-memory cache beats blocking on the network. + if _models_dev_cache: + return _serve_stale( + "Using stale in-memory models.dev cache; refreshing in background" + ) + # Stage 3: disk cache (cold-start only). A stale disk cache is + # deliberately usable so resolution doesn't hang when models.dev is + # unreachable. + disk_age = _disk_cache_age_seconds() + if disk_age is not None and (disk_data := _load_disk_cache()): + _models_dev_cache = disk_data + if disk_age >= _MODELS_DEV_CACHE_TTL: + return _serve_stale( + "Using stale models.dev disk cache (age=%.0fs); " + "refreshing in background", + disk_age, + ) + # Anchor the in-mem TTL to the file's age so an aging cache isn't + # extended by another full TTL. + _models_dev_cache_time = time.time() - disk_age + logger.debug( + "Loaded models.dev from fresh disk cache " + "(%d providers, age=%.0fs)", len(disk_data), disk_age, + ) + return _models_dev_cache + # Process-wide backoff: don't make every caller retry an unreachable + # endpoint while no usable cache exists. + if time.time() < _models_dev_retry_after: + return _models_dev_cache + # Stage 4: singleflight foreground fetch. Recheck state under the lock — # another caller may have refreshed or armed the backoff while we waited. with _models_dev_fetch_lock: - now = time.time() - if not force_refresh and (_models_dev_cache or now < _models_dev_retry_after): + if not force_refresh and (_models_dev_cache or time.time() < _models_dev_retry_after): return _models_dev_cache # Cold force_refresh: stages 1-3 were skipped, so hydrate memory from @@ -632,7 +606,6 @@ def fetch_models_dev( "Loaded stale models.dev disk cache (%d providers)", len(_models_dev_cache), ) - return _models_dev_cache @@ -646,10 +619,16 @@ def _fetch_registry(allow_network: bool) -> Dict[str, Any]: return fetch_models_dev() if allow_network else fetch_models_dev(allow_network=False) +def _registry_provider(mdev_id: str, allow_network: bool) -> Optional[Dict[str, Any]]: + """The raw models.dev provider entry, or None.""" + provider_data = _fetch_registry(allow_network).get(mdev_id) + return provider_data if isinstance(provider_data, dict) else None + + def _registry_models(mdev_id: str, *, allow_network: bool) -> Optional[Dict[str, Any]]: """The ``models`` dict of a models.dev provider entry, or None.""" - provider_data = _fetch_registry(allow_network).get(mdev_id) - if not isinstance(provider_data, dict): + provider_data = _registry_provider(mdev_id, allow_network) + if provider_data is None: return None models = provider_data.get("models", {}) return models if isinstance(models, dict) else None @@ -660,8 +639,7 @@ def _get_provider_models( ) -> Optional[Dict[str, Any]]: """Resolve a Hermes provider ID to its models dict, or None if unknown. - ``allow_network`` defaults to False — called from hot paths (vision/image - routing, capability checks) that must never block. + ``allow_network`` defaults to False — hot-path callers must never block. """ mdev_provider_id = PROVIDER_TO_MODELS_DEV.get(provider) if not mdev_provider_id: @@ -676,45 +654,31 @@ def _iter_model_entries(models: Dict[str, Any], model: str, *, suffix_fallback: The suffix fallback exists because some providers (e.g. ollama-cloud) store ``kimi-k2.6:cloud`` while the live API returns the bare name; without it context lookup falls through to stale OpenRouter metadata and - trips the 64k minimum-context guard. Every consumer shares this order so - "is this model in the catalog" means the same thing everywhere — a - suffix-keyed catalog model must count as KNOWN for ``model_overrides`` + trips the 64k minimum-context guard. Every consumer shares this order so a + suffix-keyed catalog model counts as KNOWN for ``model_overrides`` fill-gap ``_default`` semantics. """ - entry = models.get(model) - if isinstance(entry, dict): - yield model, entry - model_lower = model.lower() - for mid, mdata in models.items(): - if mid.lower() == model_lower and isinstance(mdata, dict): - yield mid, mdata - if not suffix_fallback: - return - for suffix in (":cloud", "-cloud"): - entry = models.get(model + suffix) + candidates = [model] + if suffix_fallback: + candidates += [model + suffix for suffix in (":cloud", "-cloud")] + for name in candidates: + entry = models.get(name) if isinstance(entry, dict): - yield model + suffix, entry - suffixed_lower = model_lower + suffix + yield name, entry + name_lower = name.lower() for mid, mdata in models.items(): - if mid.lower() == suffixed_lower and isinstance(mdata, dict): + if mid.lower() == name_lower and isinstance(mdata, dict): yield mid, mdata def _find_model_entry(models: Dict[str, Any], model: str) -> Optional[Dict[str, Any]]: """First catalog entry for *model* (exact, case-insensitive, suffix), or None.""" - for _mid, entry in _iter_model_entries(models, model): - return entry - return None + return next((entry for _mid, entry in _iter_model_entries(models, model)), None) def _extract_limit(entry: Any, key: str) -> Optional[int]: """Positive int ``entry.limit[key]`` or None (audio/image models have context=0).""" - if not isinstance(entry, dict): - return None - limit = entry.get("limit") - if not isinstance(limit, dict): - return None - value = limit.get(key) + value = _dict_or_empty(_dict_or_empty(entry).get("limit")).get(key) if isinstance(value, (int, float)) and value > 0: return int(value) return None @@ -734,9 +698,8 @@ def lookup_models_dev_context( entries fill the gap only when the catalog has no answer (the supported self-unblock path for models with wrong/missing context in models.dev). Catalog entries with context=0 are skipped in favour of later candidates. - - ``allow_network`` defaults to False — this runs every conversation turn - and must never block; pass True only from explicit refresh flows. + ``allow_network`` defaults to False — this runs every turn and must never + block. """ override_ctx = _override_context_window(provider, model) if override_ctx is not None: @@ -759,13 +722,12 @@ def lookup_models_dev_context( # context_window, max_output_tokens, supports_tools, supports_vision, # supports_reasoning, model_family # -# Resolution: ``model_overrides..`` is an explicit override -# that always wins over the catalog for the fields it sets (partial patch). +# ``model_overrides..`` is an explicit override that always +# wins over the catalog for the fields it sets (partial patch). # ``model_overrides.._default`` / ``model_overrides._default`` are -# FILL-GAP defaults: they apply ONLY to models the catalog does not know (the -# self-unblock path for custom/local/new models) and never displace catalog -# data — a ``_default: {context_window: 128000}`` cannot clamp every -# catalog-known model of a provider. +# FILL-GAP defaults: they apply ONLY to models the catalog does not know and +# never displace catalog data — a ``_default: {context_window: 128000}`` cannot +# clamp every catalog-known model of a provider. # # Provider keys accept the Hermes provider id or the models.dev provider id. # Model ids match exactly, then case-insensitively (mirroring catalog lookup). @@ -788,34 +750,21 @@ def _load_model_overrides() -> Dict[str, Any]: (mtime, size)-cached upstream, and an ``id(cfg)``-keyed layer here can serve stale overrides after a reload when CPython reuses the dict address. """ - try: - from hermes_cli.config import cfg_get, load_config_readonly - raw = cfg_get(load_config_readonly(), "model_overrides", default={}) - return raw if isinstance(raw, dict) else {} - except Exception: - return {} + return _dict_or_empty(_cfg_get("model_overrides", default={})) def _provider_override_section(provider: str) -> Optional[Dict[str, Any]]: """Override section for *provider* (keyed by Hermes OR models.dev id), or None.""" overrides = _load_model_overrides() - if not overrides: - return None provider_key = (provider or "").strip() - if not provider_key: + if not overrides or not provider_key: return None - - candidates = [provider_key] - mapped = PROVIDER_TO_MODELS_DEV.get(provider_key) - if mapped and mapped != provider_key: - candidates.append(mapped) - # Reverse: caller passed a models.dev id, config keyed by Hermes id. - for hermes_id in _models_dev_to_hermes_ids(provider_key): - if hermes_id != provider_key: - candidates.append(hermes_id) - + # Forward (Hermes → models.dev id) and reverse (caller passed a models.dev + # id, config keyed by Hermes id) aliases. + candidates = [provider_key, PROVIDER_TO_MODELS_DEV.get(provider_key)] + candidates += _models_dev_to_hermes_ids(provider_key) for key in candidates: - section = overrides.get(key) + section = overrides.get(key) if key else None if isinstance(section, dict): return section return None @@ -825,21 +774,15 @@ def _explicit_model_override(provider: str, model: str) -> Optional[Dict[str, An """Explicit per-provider+model override dict (exact, then case-insensitive skipping the ``_default`` sentinel), or None.""" model_key = (model or "").strip() - if not model_key: - return None - section = _provider_override_section(provider) + section = _provider_override_section(provider) if model_key else None if section is None: return None - entry = section.get(model_key) if isinstance(entry, dict): return entry - model_lower = model_key.lower() for mid, mdata in section.items(): - if mid == "_default": - continue - if mid.lower() == model_lower and isinstance(mdata, dict): + if mid != "_default" and mid.lower() == model_lower and isinstance(mdata, dict): return mdata return None @@ -847,14 +790,10 @@ def _explicit_model_override(provider: str, model: str) -> Optional[Dict[str, An def _default_model_override(provider: str) -> Optional[Dict[str, Any]]: """Fill-gap ``_default`` override: per-provider first, then global; or None.""" section = _provider_override_section(provider) - if section is not None: - default = section.get("_default") - if isinstance(default, dict): - return default + if section is not None and isinstance(section.get("_default"), dict): + return section["_default"] global_default = _load_model_overrides().get("_default") - if isinstance(global_default, dict): - return global_default - return None + return global_default if isinstance(global_default, dict) else None def _override_for( @@ -862,10 +801,8 @@ def _override_for( ) -> Optional[Dict[str, Any]]: """Explicit override if any; else the ``_default`` only on a catalog miss.""" explicit = _explicit_model_override(provider, model) - if explicit is not None: + if explicit is not None or catalog_hit: return explicit - if catalog_hit: - return None return _default_model_override(provider) @@ -913,29 +850,23 @@ def _override_to_catalog_shape( ) -> Tuple[Dict[str, Any], Optional[bool]]: """Translate canonical override keys into a models.dev-shaped patch. - Consumers read the raw catalog shape (``limit.context``, ``tool_call``, ...) - while users write ONE canonical schema, so this boundary translates. Returns ``(patch, vision)`` — vision is out-of-band because it maps onto the ``modalities.input`` list rather than a scalar field. """ patch: Dict[str, Any] = {} - limit: Dict[str, Any] = {} - ctx = _override_int(override, "context_window") - if ctx is not None: - limit["context"] = ctx - out = _override_int(override, "max_output_tokens") - if out is not None: - limit["output"] = out + limit = { + catalog_key: value + for catalog_key, override_key in (("context", "context_window"), ("output", "max_output_tokens")) + if (value := _override_int(override, override_key)) is not None + } if limit: patch["limit"] = limit - if "supports_tools" in override: - patch["tool_call"] = bool(override["supports_tools"]) - if "supports_reasoning" in override: - patch["reasoning"] = bool(override["supports_reasoning"]) + for override_key, catalog_key in (("supports_tools", "tool_call"), ("supports_reasoning", "reasoning")): + if override_key in override: + patch[catalog_key] = bool(override[override_key]) vision: Optional[bool] = None if "supports_vision" in override: - vision = bool(override["supports_vision"]) - patch["attachment"] = vision + vision = patch["attachment"] = bool(override["supports_vision"]) if "model_family" in override: patch["family"] = str(override["model_family"] or "") return patch, vision @@ -951,13 +882,9 @@ def _merge_catalog_entry_with_override( merged = dict(raw) limit_patch = shaped.pop("limit", None) if limit_patch: - base_limit = raw.get("limit") - base_limit = dict(base_limit) if isinstance(base_limit, dict) else {} - base_limit.update(limit_patch) - merged["limit"] = base_limit + merged["limit"] = {**_dict_or_empty(raw.get("limit")), **limit_patch} if vision_override is not None: - base_mods = raw.get("modalities") - base_mods = dict(base_mods) if isinstance(base_mods, dict) else {} + base_mods = dict(_dict_or_empty(raw.get("modalities"))) input_mods = base_mods.get("input") input_mods = list(input_mods) if isinstance(input_mods, list) else [] if vision_override and "image" not in input_mods: @@ -970,6 +897,22 @@ def _merge_catalog_entry_with_override( return merged +def _apply_overrides( + provider: str, model: str, entry: Optional[Dict[str, Any]] +) -> Optional[Dict[str, Any]]: + """Catalog *entry* patched by its override; ``_UNKNOWN_MODEL_BASE`` patched + by a fill-gap override on a catalog miss; None when neither exists. + + The override is selected AFTER the catalog lookup: _default only fills misses. + """ + override = _override_for(provider, model, catalog_hit=entry is not None) + if override is None: + return entry + return _merge_catalog_entry_with_override( + entry if entry is not None else _UNKNOWN_MODEL_BASE, override + ) + + # --------------------------------------------------------------------------- # Model capability metadata # --------------------------------------------------------------------------- @@ -978,8 +921,7 @@ def _entry_supports_vision(entry: Dict[str, Any]) -> bool: """Prefer explicit ``modalities.input`` (the older ``attachment`` flag can be stale or too broad for image routing); fall back to it only when the input modalities are absent/invalid.""" - input_mods = entry.get("modalities", {}) - input_mods = input_mods.get("input") if isinstance(input_mods, dict) else None + input_mods = _dict_or_empty(entry.get("modalities", {})).get("input") if isinstance(input_mods, list): return "image" in input_mods return bool(entry.get("attachment", False)) @@ -994,21 +936,13 @@ def get_model_capabilities( they set; ``_default`` entries fill the gap only for models the catalog does not know. Unspecified fields fall through to the catalog value, or to safe defaults (tools on, vision/reasoning off, 200K/8K) when absent. - ``allow_network`` defaults to False — vision/image routing is a hot path. """ models = _get_provider_models(provider, allow_network=allow_network) entry = _find_model_entry(models, model) if models is not None else None - - # Select the override AFTER the catalog lookup: _default only fills misses. - override = _override_for(provider, model, catalog_hit=entry is not None) - if entry is None and override is None: + raw = _apply_overrides(provider, model, entry) + if raw is None: return None - - raw = entry if entry is not None else _UNKNOWN_MODEL_BASE - if override is not None: - raw = _merge_catalog_entry_with_override(raw, override) - return ModelCapabilities( supports_tools=bool(raw.get("tool_call", False)), supports_vision=_entry_supports_vision(raw), @@ -1102,13 +1036,8 @@ def list_agentic_models( def _parse_model_info(model_id: str, raw: Dict[str, Any], provider_id: str) -> ModelInfo: """Convert a raw models.dev model entry dict into a ModelInfo dataclass.""" - cost = raw.get("cost") or {} - if not isinstance(cost, dict): - cost = {} - - modalities = raw.get("modalities") or {} - if not isinstance(modalities, dict): - modalities = {} + cost = _dict_or_empty(raw.get("cost")) + modalities = _dict_or_empty(raw.get("modalities")) input_mods = modalities.get("input") or [] output_mods = modalities.get("output") or [] @@ -1165,10 +1094,8 @@ def get_provider_info( setup (``resolve_provider_full``). Hot-path callers pass False. """ mdev_id = PROVIDER_TO_MODELS_DEV.get(provider_id, provider_id) - raw = _fetch_registry(allow_network).get(mdev_id) - if not isinstance(raw, dict): - return None - return _parse_provider_info(mdev_id, raw) + raw = _registry_provider(mdev_id, allow_network) + return _parse_provider_info(mdev_id, raw) if raw is not None else None def get_model_info( @@ -1177,26 +1104,16 @@ def get_model_info( """Full model metadata by Hermes or models.dev provider ID (exact match, then case-insensitive), or None if not found. - ``model_overrides`` use the same canonical schema as every other consumer - and are translated into the catalog shape here with sub-dicts merged, not - clobbered. EXPLICIT entries patch known catalog models; ``_default`` - entries fill the gap only for models the catalog does not know. - + ``model_overrides`` use the same canonical schema as every other consumer: + EXPLICIT entries patch known catalog models; ``_default`` entries fill the + gap only for models the catalog does not know. ``allow_network`` defaults to False — cost guard and inventory are hot paths. """ mdev_id = PROVIDER_TO_MODELS_DEV.get(provider_id, provider_id) - - def _resolve(mid: str, raw: Dict[str, Any], *, catalog_hit: bool) -> Optional[ModelInfo]: - override = _override_for(provider_id, model_id, catalog_hit=catalog_hit) - if override is not None: - raw = _merge_catalog_entry_with_override(raw, override) - elif not catalog_hit: - return None - return _parse_model_info(mid, raw, mdev_id) - models = _registry_models(mdev_id, allow_network=allow_network) + mid, entry = model_id, None if models is not None: - for mid, raw in _iter_model_entries(models, model_id, suffix_fallback=False): - return _resolve(mid, raw, catalog_hit=True) + mid, entry = next(_iter_model_entries(models, model_id, suffix_fallback=False), (model_id, None)) # Not in catalog — an override (explicit or _default) may still provide it. - return _resolve(model_id, _UNKNOWN_MODEL_BASE, catalog_hit=False) + raw = _apply_overrides(provider_id, model_id, entry) + return _parse_model_info(mid, raw, mdev_id) if raw is not None else None