diff --git a/agent/model_metadata.py b/agent/model_metadata.py index 38f49b041a..1777e34739 100644 --- a/agent/model_metadata.py +++ b/agent/model_metadata.py @@ -13,7 +13,7 @@ import os import re import time from pathlib import Path -from typing import Any, Dict, List, Optional, Tuple, TYPE_CHECKING +from typing import Any, Callable, Dict, List, Optional, Tuple, TYPE_CHECKING from urllib.parse import urlparse import yaml @@ -71,6 +71,7 @@ def _resolve_requests_verify(base_url: str = "") -> bool | str: return val return True + # Compatibility snapshot for callers that inspect this private constant. # Prefix routing below queries the registry live so later registrations work. try: @@ -85,13 +86,11 @@ _PROVIDER_PREFIXES: frozenset[str] = frozenset( for value in (profile.name, *profile.aliases) ) - _OLLAMA_TAG_PATTERN = re.compile( r"^(\d+\.?\d*b|latest|stable|q\d|fp?\d|instruct|chat|coder|vision|text)", re.IGNORECASE, ) - # Tailscale CGNAT (RFC 6598): `ipaddress.is_private` excludes it, yet Ollama # reached over Tailscale must count as local (timeout auto-bumps). _TAILSCALE_CGNAT = ipaddress.IPv4Network("100.64.0.0/10") @@ -106,52 +105,52 @@ def _strip_provider_prefix(model: str) -> str: if ":" not in model or model.startswith("http"): return model prefix, suffix = model.split(":", 1) - prefix_lower = prefix.strip().lower() try: from providers import get_provider_profile - - is_provider = get_provider_profile(prefix_lower) is not None + is_provider = get_provider_profile(prefix.strip().lower()) is not None except Exception: is_provider = False - if is_provider: - # Don't strip if suffix looks like an Ollama tag (e.g. "7b", "latest", "q4_0") - if _OLLAMA_TAG_PATTERN.match(suffix.strip()): - return model + if is_provider and not _OLLAMA_TAG_PATTERN.match(suffix.strip()): return suffix return model + _model_metadata_cache: Dict[str, Dict[str, Any]] = {} _model_metadata_cache_time: float = 0 -_novita_metadata_cache: Dict[str, Dict[str, Any]] = {} -_novita_metadata_cache_time: float = 0 _MODEL_CACHE_TTL = 3600 _endpoint_model_metadata_cache: Dict[str, Dict[str, Dict[str, Any]]] = {} _endpoint_model_metadata_cache_time: Dict[str, float] = {} _ENDPOINT_MODEL_CACHE_TTL = 300 # Server-type verdicts: (server_type, monotonic_ts). Positive verdicts live an # hour (so a server swap on the same port is eventually re-detected); a None -# verdict gets the short TTL so a transient failure (server starting, key being -# fixed) recovers in minutes while still not re-running the waterfall each turn. +# verdict gets the short TTL so a transient failure recovers in minutes while +# still not re-running the waterfall each turn. _ENDPOINT_PROBE_TTL_SECONDS = 3600.0 _ENDPOINT_PROBE_FAILURE_TTL_SECONDS = 300.0 _endpoint_probe_path_cache: Dict[str, tuple] = {} # Routable-but-dead endpoints (corp LAN off-VPN) blackhole TCP: every probe -# waits out its full connect timeout and startup stalls for a minute. Once ANY -# probe observed a connect timeout, later probes short-circuit for a while. -# Pure bookkeeping — no network I/O, fires only after a real timeout was paid. +# waits out its full connect timeout. Once ANY probe observed a connect +# timeout, later probes short-circuit for a while. Fires only after a real +# timeout was paid. _ENDPOINT_BLACKHOLE_TTL_SECONDS = 30.0 _endpoint_blackhole_cache: Dict[str, float] = {} # host:port -> monotonic ts -def _endpoint_host_key(base_url: str) -> Optional[str]: - """``host:port`` key (None without a host) so every probe path for one server shares an entry.""" +def _parse_base_url(base_url: str, scheme: str = "http"): + """``urlparse`` of the normalized URL (``scheme`` prepended when absent); None when empty.""" normalized = _normalize_base_url(base_url) if not normalized: return None - url = normalized if "://" in normalized else f"http://{normalized}" + return urlparse(normalized if "://" in normalized else f"{scheme}://{normalized}") + + +def _endpoint_host_key(base_url: str) -> Optional[str]: + """``host:port`` key (None without a host) so every probe path for one server shares an entry.""" try: - parsed = urlparse(url) + parsed = _parse_base_url(base_url) + if parsed is None: + return None host = parsed.hostname port = parsed.port or (443 if parsed.scheme == "https" else 80) except Exception: @@ -176,9 +175,7 @@ def _endpoint_blackholed(base_url: str) -> bool: if _ENDPOINT_BLACKHOLE_TTL_SECONDS <= 0: return False key = _endpoint_host_key(base_url) - if key is None: - return False - seen = _endpoint_blackhole_cache.get(key) + seen = _endpoint_blackhole_cache.get(key) if key is not None else None if seen is None: return False if (time.monotonic() - seen) >= _ENDPOINT_BLACKHOLE_TTL_SECONDS: @@ -187,6 +184,12 @@ def _endpoint_blackholed(base_url: str) -> bool: return True +def _note_if_connect_timeout(exc: BaseException, base_url: str) -> None: + """Blackhole ``base_url`` when ``exc`` is a connect-phase timeout.""" + if _is_connect_timeout(exc): + _note_endpoint_blackholed(base_url) + + def _is_connect_timeout(exc: BaseException) -> bool: """True for connect-phase timeouts raised by httpx or requests. @@ -201,11 +204,10 @@ def _is_connect_timeout(exc: BaseException) -> bool: pass try: from requests.exceptions import ConnectTimeout - if isinstance(exc, ConnectTimeout): - return True + return isinstance(exc, ConnectTimeout) except Exception: - pass - return False + return False + # Disk L2 for local-endpoint probes so back-to-back CLI cold starts skip the # waterfall. Only SUCCESSFUL probes are persisted (a down server must not pin @@ -278,54 +280,39 @@ def _local_probe_disk_put(kind: str, key: str, value: Any) -> None: def _get_model_metadata_cache_path() -> Path: - """Return path to the OpenRouter model metadata disk cache.""" + """Path to the OpenRouter model metadata disk cache.""" return _cache_file("openrouter_model_metadata.json") def _model_metadata_disk_cache_age_seconds() -> Optional[float]: - """Return disk-cache age in seconds, or None if freshness is unknown.""" + """Disk-cache age in seconds, or None if freshness is unknown.""" try: - cache_path = _get_model_metadata_cache_path() - if not cache_path.exists(): - return None - age = time.time() - cache_path.stat().st_mtime - if age < 0: - return None - return age + age = time.time() - _get_model_metadata_cache_path().stat().st_mtime + return age if age >= 0 else None except Exception: return None def _load_model_metadata_disk_cache() -> Dict[str, Dict[str, Any]]: - """Load processed OpenRouter metadata cache from disk.""" + """Processed OpenRouter metadata cache from disk ({} on any failure).""" try: - cache_path = _get_model_metadata_cache_path() - with cache_path.open("r", encoding="utf-8") as f: + with _get_model_metadata_cache_path().open("r", encoding="utf-8") as f: data = json.load(f) if not isinstance(data, dict): return {} - return { - str(key): value - for key, value in data.items() - if isinstance(value, dict) - } + return {str(key): value for key, value in data.items() if isinstance(value, dict)} except Exception as e: logger.debug("Failed to load OpenRouter model metadata disk cache: %s", e) return {} def _save_model_metadata_disk_cache(data: Dict[str, Dict[str, Any]]) -> None: - """Save processed OpenRouter metadata cache to disk atomically.""" try: - atomic_json_write( - _get_model_metadata_cache_path(), - data, - indent=0, - separators=(",", ":"), - ) + atomic_json_write(_get_model_metadata_cache_path(), data, indent=0, separators=(",", ":")) except Exception as e: logger.debug("Failed to save OpenRouter model metadata disk cache: %s", e) + def _get_endpoint_metadata_cache_path() -> Path: """On-disk memo of remote ``/models`` probes (see ``_endpoint_disk_cache_get``).""" return _cache_file("endpoint_model_metadata.json") @@ -372,9 +359,8 @@ _FALLBACK_WARNED: set = set() def _warn_context_length_fallback(model: str, base_url: str) -> None: - """Warn (once per model+endpoint) that context detection failed and the - hard default is being used, so small-context models (8K, 32K) don't - silently get 256K and cause hard-to-debug API failures.""" + """Warn once per model+endpoint that detection failed and the hard default is used, + so small-context models (8K, 32K) don't silently get 256K and fail at the API.""" key = (model, base_url or "") if key in _FALLBACK_WARNED: return @@ -386,6 +372,7 @@ def _warn_context_length_fallback(model: str, base_url: str) -> None: model, base_url or "default", f"{DEFAULT_FALLBACK_CONTEXT:,}", ) + # Sessions, model switches and cron jobs reject models below this: too little # working memory for tool-calling workflows. MINIMUM_CONTEXT_LENGTH = 64_000 @@ -531,8 +518,7 @@ DEFAULT_CONTEXT_LENGTHS = { # xAI Grok models that ACCEPT `reasoning.effort` (verified live against # /v1/responses). Unlisted Grok models still reason natively but 400 on the -# parameter ("Model X does not support parameter reasoningEffort"), so callers -# must send no `reasoning` key rather than a default `medium`. +# parameter, so callers must send no `reasoning` key rather than a default `medium`. _GROK_EFFORT_CAPABLE_PREFIXES = ( "grok-3-mini", "grok-4.20-multi-agent", @@ -544,19 +530,13 @@ _GROK_EFFORT_CAPABLE_PREFIXES = ( def grok_supports_reasoning_effort(model: str) -> bool: """Allowlist check (aggregator prefixes like ``x-ai/`` stripped); unknown Grok models get no effort dial.""" - name = (model or "").strip().lower() - if not name: - return False - if "/" in name: - name = name.rsplit("/", 1)[-1] - return any(name.startswith(prefix) for prefix in _GROK_EFFORT_CAPABLE_PREFIXES) + name = (model or "").strip().lower().rsplit("/", 1)[-1] + return bool(name) and any(name.startswith(prefix) for prefix in _GROK_EFFORT_CAPABLE_PREFIXES) def is_grok_46_family(model: str) -> bool: - """Return whether *model* is a Grok 4.6 family identifier.""" - name = (model or "").strip().lower().replace("_", "-") - if "/" in name: - name = name.rsplit("/", 1)[-1] + """Whether *model* is a Grok 4.6 family identifier.""" + name = (model or "").strip().lower().replace("_", "-").rsplit("/", 1)[-1] return name == "grok-4.6" or name.startswith("grok-4.6-") @@ -597,9 +577,7 @@ def _normalize_base_url(base_url: str) -> str: def _auth_headers(api_key: str = "") -> Dict[str, str]: token = str(api_key or "").strip() - if not token: - return {} - return {"Authorization": f"Bearer {token}"} + return {"Authorization": f"Bearer {token}"} if token else {} def _is_openrouter_base_url(base_url: str) -> bool: @@ -662,10 +640,9 @@ except Exception: def _infer_provider_from_url(base_url: str) -> Optional[str]: """models.dev provider name for a base URL (custom endpoints need no explicit provider).""" - normalized = _normalize_base_url(base_url) - if not normalized: + parsed = _parse_base_url(base_url, "https") + if parsed is None: return None - parsed = urlparse(normalized if "://" in normalized else f"https://{normalized}") host = parsed.netloc.lower() or parsed.path.lower() for url_part, provider in _URL_TO_PROVIDER.items(): if url_part in host: @@ -674,12 +651,11 @@ def _infer_provider_from_url(base_url: str) -> Optional[str]: def _lmstudio_server_root(base_url: str) -> str: - """Return the LM Studio server root for native ``/api/v1`` endpoints.""" - root = _normalize_base_url(base_url).rstrip("/") + """LM Studio server root for native ``/api/v1`` endpoints.""" + root = _normalize_base_url(base_url) for suffix in ("/api/v1", "/api", "/v1"): if root.endswith(suffix): - root = root[: -len(suffix)].rstrip("/") - break + return root[: -len(suffix)].rstrip("/") return root @@ -690,9 +666,7 @@ def _is_known_provider_base_url(base_url: str) -> bool: def _server_root(base_url: str) -> str: """Probe root for a local server: IPv4-resolved, ``/v1`` suffix stripped.""" server_url = _localhost_to_ipv4(base_url.rstrip("/")) - if server_url.endswith("/v1"): - server_url = server_url[:-3] - return server_url + return server_url[:-3] if server_url.endswith("/v1") else server_url def _longest_key_match(table: Dict[str, int], model_lower: str) -> Optional[Tuple[str, int]]: @@ -713,28 +687,28 @@ def _ollama_show_context(data: Dict[str, Any], *, gguf_first: bool, minimum: Opt ``parameters`` -> ``num_ctx`` is the Modelfile override (the RUNTIME window Ollama allocates KV cache for); ``model_info.*.context_length`` is the GGUF training max, which can exceed num_ctx. Local users control num_ctx, so - local probes prefer it; hosted Ollama operators may cap num_ctx - arbitrarily, so hosted probes prefer the GGUF value (``gguf_first``). + local probes prefer it; hosted operators may cap num_ctx arbitrarily, so + hosted probes prefer the GGUF value (``gguf_first``). """ + def _ok(ctx: int) -> bool: + return minimum is None or ctx >= minimum + def _num_ctx() -> Optional[int]: for line in data.get("parameters", "").split("\n"): - if "num_ctx" in line: - parts = line.strip().split() - if len(parts) >= 2: - try: - ctx = int(parts[-1]) - except ValueError: - continue - if minimum is None or ctx >= minimum: - return ctx + parts = line.strip().split() + if "num_ctx" in line and len(parts) >= 2: + try: + ctx = int(parts[-1]) + except ValueError: + continue + if _ok(ctx): + return ctx return None def _gguf() -> Optional[int]: for key, value in data.get("model_info", {}).items(): - if "context_length" in key and isinstance(value, (int, float)): - ctx = int(value) - if minimum is None or ctx >= minimum: - return ctx + if "context_length" in key and isinstance(value, (int, float)) and _ok(int(value)): + return int(value) return None for reader in ((_gguf, _num_ctx) if gguf_first else (_num_ctx, _gguf)): @@ -759,9 +733,8 @@ def _endpoint_scoped_context_length(model: str, base_url: str) -> Optional[int]: not. NVIDIA NIM serves deepseek-v4-pro at 262,144 while DeepSeek's native endpoint is 1M; the lower limit stays scoped to NVIDIA. """ - normalized = _normalize_base_url(base_url) try: - parsed = urlparse(normalized) + parsed = urlparse(_normalize_base_url(base_url)) port = parsed.port except ValueError: return None @@ -794,11 +767,13 @@ def _skip_persistent_context_cache(base_url: str, provider: str) -> bool: return (provider or "").strip().lower() in {"lmstudio", "openai-codex"} -def _maybe_cache_local_context_length( - model: str, - base_url: str, - length: int, -) -> None: +def _save_unless_skipped(model: str, base_url: str, ctx: int, provider: str) -> None: + """Persist ``ctx`` unless the provider opts out of the disk cache.""" + if not _skip_persistent_context_cache(base_url, provider): + save_context_length(model, base_url, ctx) + + +def _maybe_cache_local_context_length(model: str, base_url: str, length: int) -> None: """Persist a probed local window only at/above MINIMUM_CONTEXT_LENGTH. Sub-minimum windows are still returned so agent_init can reject them, but @@ -818,12 +793,7 @@ def _probe_local_context_length(model: str, base_url: str, api_key: str, provide return None -def _reconcile_local_cached_context_length( - model: str, - base_url: str, - cached: int, - api_key: str = "", -) -> int: +def _reconcile_local_cached_context_length(model: str, base_url: str, cached: int, api_key: str = "") -> int: """Return *cached* unless a live local probe reports a different limit. Operators restart vLLM/Ollama with a new --max-model-len / num_ctx under @@ -831,48 +801,42 @@ def _reconcile_local_cached_context_length( probe keeps it. Sub-minimum live windows invalidate but are not persisted. """ live_ctx = _query_local_context_length(model, base_url, api_key=api_key) - if live_ctx and live_ctx > 0 and live_ctx != cached: - if live_ctx < MINIMUM_CONTEXT_LENGTH: - logger.info( - "Live local probe for %s@%s reports %s (< minimum %s); " - "invalidating stale cache — agent init should reject", - model, base_url, f"{live_ctx:,}", f"{MINIMUM_CONTEXT_LENGTH:,}", - ) - _invalidate_cached_context_length(model, base_url) - return live_ctx + if not (live_ctx and live_ctx > 0 and live_ctx != cached): + return cached + if live_ctx < MINIMUM_CONTEXT_LENGTH: logger.info( - "Reconciling stale local cache entry %s@%s: %s -> %s (live probe)", - model, base_url, f"{cached:,}", f"{live_ctx:,}", + "Live local probe for %s@%s reports %s (< minimum %s); " + "invalidating stale cache — agent init should reject", + model, base_url, f"{live_ctx:,}", f"{MINIMUM_CONTEXT_LENGTH:,}", ) _invalidate_cached_context_length(model, base_url) - _maybe_cache_local_context_length(model, base_url, live_ctx) return live_ctx - return cached + logger.info( + "Reconciling stale local cache entry %s@%s: %s -> %s (live probe)", + model, base_url, f"{cached:,}", f"{live_ctx:,}", + ) + _invalidate_cached_context_length(model, base_url) + _maybe_cache_local_context_length(model, base_url, live_ctx) + return live_ctx def is_local_endpoint(base_url: str) -> bool: """True for loopback, container-internal DNS, unqualified hosts, RFC-1918, link-local and Tailscale CGNAT (so a trusted Ollama box over Tailscale gets the same timeout auto-bumps as localhost).""" - normalized = _normalize_base_url(base_url) - if not normalized: - return False - url = normalized if "://" in normalized else f"http://{normalized}" try: - parsed = urlparse(url) + parsed = _parse_base_url(base_url) + if parsed is None: + return False host = parsed.hostname or "" except Exception: return False - if host in _LOCAL_HOSTS: - return True - # Docker / Podman / Lima internal DNS names (e.g. host.docker.internal) - if any(host.endswith(suffix) for suffix in _CONTAINER_LOCAL_SUFFIXES): + if host in _LOCAL_HOSTS or any(host.endswith(suffix) for suffix in _CONTAINER_LOCAL_SUFFIXES): return True # Unqualified hostnames (no dots) are local by definition — Docker # Compose service names, /etc/hosts entries, or mDNS names. if host and "." not in host: return True - # RFC-1918 private ranges, link-local, and Tailscale CGNAT try: addr = ipaddress.ip_address(host) if addr.is_private or addr.is_loopback or addr.is_link_local: @@ -881,21 +845,21 @@ def is_local_endpoint(base_url: str) -> bool: return True except ValueError: pass - # Bare IP that looks like a private range (e.g. 172.26.x.x for WSL) - # or Tailscale CGNAT (100.64.x.x–100.127.x.x). + # Dotted quad that ipaddress rejected but still looks like a private range + # (e.g. 172.26.x.x for WSL) or Tailscale CGNAT (100.64.x.x–100.127.x.x). parts = host.split(".") - if len(parts) == 4: - try: - first, second = int(parts[0]), int(parts[1]) - except ValueError: - return False - return ( - first == 10 - or (first == 172 and 16 <= second <= 31) - or (first == 192 and second == 168) - or (first == 100 and 64 <= second <= 127) - ) - return False + if len(parts) != 4: + return False + try: + first, second = int(parts[0]), int(parts[1]) + except ValueError: + return False + return ( + first == 10 + or (first == 172 and 16 <= second <= 31) + or (first == 192 and second == 168) + or (first == 100 and 64 <= second <= 127) + ) def _localhost_to_ipv4(url: str) -> str: @@ -907,12 +871,7 @@ def _localhost_to_ipv4(url: str) -> str: """ if not url or not isinstance(url, str): return url # non-string values (test doubles, lazy config) pass through - return re.sub( - r"^(https?://)localhost(?=[:/]|$)", - r"\g<1>127.0.0.1", - url, - count=1, - ) + return re.sub(r"^(https?://)localhost(?=[:/]|$)", r"\g<1>127.0.0.1", url, count=1) def detect_local_server_type(base_url: str, api_key: str = "") -> Optional[str]: @@ -927,11 +886,7 @@ def detect_local_server_type(base_url: str, api_key: str = "") -> Optional[str]: cached = _endpoint_probe_path_cache.get(server_url) if cached is not None: - ttl = ( - _ENDPOINT_PROBE_TTL_SECONDS - if cached[0] is not None - else _ENDPOINT_PROBE_FAILURE_TTL_SECONDS - ) + ttl = _ENDPOINT_PROBE_TTL_SECONDS if cached[0] is not None else _ENDPOINT_PROBE_FAILURE_TTL_SECONDS if (time.monotonic() - cached[1]) < ttl: return cached[0] @@ -945,14 +900,6 @@ def detect_local_server_type(base_url: str, api_key: str = "") -> Optional[str]: _endpoint_probe_path_cache[server_url] = (disk_hit, time.monotonic()) return disk_hit - headers = _auth_headers(api_key) - - def _probe_failed(exc: Exception) -> None: - """Swallow a probe error; on a connect timeout re-raise so the remaining legs are skipped.""" - if _is_connect_timeout(exc): - _note_endpoint_blackholed(server_url) - raise exc - def _lm_studio(client) -> bool: return client.get(f"{lmstudio_url}/api/v1/models").status_code == 200 @@ -976,23 +923,24 @@ def detect_local_server_type(base_url: str, api_key: str = "") -> Optional[str]: waterfall = (("lm-studio", _lm_studio), ("ollama", _ollama), ("llamacpp", _llamacpp), ("vllm", _vllm)) result: Optional[str] = None try: - with httpx.Client(timeout=2.0, headers=headers) as client: + with httpx.Client(timeout=2.0, headers=_auth_headers(api_key)) as client: for name, probe in waterfall: try: if probe(client): result = name break except Exception as exc: - _probe_failed(exc) + # A connect timeout condemns the host: skip the remaining legs. + if _is_connect_timeout(exc): + _note_endpoint_blackholed(server_url) + raise except Exception: pass + # Negative verdict in memory only (never on disk — failures are often transient). + _endpoint_probe_path_cache[server_url] = (result, time.monotonic()) if result is not None: - _endpoint_probe_path_cache[server_url] = (result, time.monotonic()) _local_probe_disk_put("server_type", server_url, result) - else: - # Negative verdict in memory only (never on disk — failures are often transient). - _endpoint_probe_path_cache[server_url] = (None, time.monotonic()) return result @@ -1015,20 +963,17 @@ def _coerce_reasonable_int(value: Any, minimum: int = 1024, maximum: int = 10_00 result = int(value) except (TypeError, ValueError): return None - if minimum <= result <= maximum: - return result - return None + return result if minimum <= result <= maximum else None def _extract_first_int(payload: Dict[str, Any], keys: tuple[str, ...]) -> Optional[int]: keyset = {key.lower() for key in keys} for mapping in _iter_nested_dicts(payload): for key, value in mapping.items(): - if str(key).lower() not in keyset: - continue - coerced = _coerce_reasonable_int(value) - if coerced is not None: - return coerced + if str(key).lower() in keyset: + coerced = _coerce_reasonable_int(value) + if coerced is not None: + return coerced return None @@ -1064,10 +1009,8 @@ def _context_length_from_model_payload(payload: Dict[str, Any]) -> Optional[int] if ctx is not None: return ctx raw = payload.get("max_tokens") - if isinstance(raw, (int, float)): - ivalue = int(raw) - if ivalue > 0: - return ivalue + if isinstance(raw, (int, float)) and int(raw) > 0: + return int(raw) return None @@ -1086,8 +1029,8 @@ def _extract_pricing(payload: Dict[str, Any]) -> Dict[str, Any]: return _per_token(payload, novita_fields, lambda v: v / 10_000 / 1_000_000) # DeepInfra ships pricing under ``metadata.pricing`` in $/MTok. - metadata = payload.get("metadata") if isinstance(payload.get("metadata"), dict) else None - deepinfra_pricing = metadata.get("pricing") if metadata else None + metadata = payload.get("metadata") + deepinfra_pricing = metadata.get("pricing") if isinstance(metadata, dict) else None deepinfra_fields = {"prompt": "input_tokens", "completion": "output_tokens", "cache_read": "cache_read_tokens"} if isinstance(deepinfra_pricing, dict) and any(k in deepinfra_pricing for k in deepinfra_fields.values()): return _per_token(deepinfra_pricing, deepinfra_fields, lambda v: v / 1_000_000) @@ -1101,8 +1044,6 @@ def _extract_pricing(payload: Dict[str, Any]) -> Dict[str, Any]: } for mapping in _iter_nested_dicts(payload): normalized = {str(key).lower(): value for key, value in mapping.items()} - if not any(any(alias in normalized for alias in aliases) for aliases in alias_map.values()): - continue pricing: Dict[str, Any] = {} for target, aliases in alias_map.items(): for alias in aliases: @@ -1117,8 +1058,7 @@ def _extract_pricing(payload: Dict[str, Any]) -> Dict[str, Any]: def _add_model_aliases(cache: Dict[str, Dict[str, Any]], model_id: str, entry: Dict[str, Any]) -> None: cache[model_id] = entry if "/" in model_id: - bare_model = model_id.split("/", 1)[1] - cache.setdefault(bare_model, entry) + cache.setdefault(model_id.split("/", 1)[1], entry) def fetch_model_metadata(force_refresh: bool = False) -> Dict[str, Dict[str, Any]]: @@ -1143,10 +1083,8 @@ def fetch_model_metadata(force_refresh: bool = False) -> Dict[str, Dict[str, Any # stage through proxies that 403 CONNECT, ballooning to minutes. response = requests.get(OPENROUTER_MODELS_URL, timeout=(5, 10), verify=_resolve_requests_verify()) response.raise_for_status() - data = response.json() - cache = {} - for model in data.get("data", []): + for model in response.json().get("data", []): model_id = model.get("id", "") entry = { "context_length": model.get("context_length", 128000), @@ -1195,6 +1133,16 @@ def _endpoint_model_entry(model: Dict[str, Any], model_id: str, context_length: return entry +def _lmstudio_loaded_context(model: Dict[str, Any]) -> Optional[int]: + """Context of the first loaded LM Studio instance (the runtime value), else None.""" + for inst in model.get("loaded_instances", []) or []: + cfg = inst.get("config", {}) if isinstance(inst, dict) else None + ctx = cfg.get("context_length") if isinstance(cfg, dict) else None + if isinstance(ctx, int) and ctx > 0: + return ctx + return None + + def _lmstudio_native_models(normalized: str, headers: Dict[str, str]) -> Dict[str, Dict[str, Any]]: """LM Studio ``/api/v1/models`` → cache; context comes from the first loaded instance.""" response = requests.get( @@ -1211,16 +1159,7 @@ def _lmstudio_native_models(normalized: str, headers: Dict[str, str]) -> Dict[st model_id = model.get("key") or model.get("id") if not model_id: continue - context_length = None - for inst in model.get("loaded_instances", []) or []: - if not isinstance(inst, dict): - continue - cfg = inst.get("config", {}) - ctx = cfg.get("context_length") if isinstance(cfg, dict) else None - if isinstance(ctx, int) and ctx > 0: - context_length = ctx - break - entry = _endpoint_model_entry(model, model_id, context_length) + entry = _endpoint_model_entry(model, model_id, _lmstudio_loaded_context(model)) _add_model_aliases(cache, model_id, entry) alt_id = model.get("id") if isinstance(alt_id, str) and alt_id and alt_id != model_id: @@ -1243,10 +1182,13 @@ def _apply_llamacpp_props(cache: Dict[str, Dict[str, Any]], request_candidate: s resp = requests.get(base + "/props", params=params, headers=headers, timeout=5, verify=verify) return resp + def _n_ctx(props: Dict[str, Any]) -> Any: + return (props.get("default_generation_settings") or {}).get("n_ctx") + props_resp = _props() if props_resp.ok: props = props_resp.json() - n_ctx = props.get("default_generation_settings", {}).get("n_ctx") + n_ctx = _n_ctx(props) model_alias = props.get("model_alias", "") if n_ctx and model_alias and model_alias in cache: cache[model_alias]["context_length"] = n_ctx @@ -1263,11 +1205,26 @@ def _apply_llamacpp_props(cache: Dict[str, Dict[str, Any]], request_candidate: s continue pr = _props({"model": child_id}) if pr.ok: - child_ctx = (pr.json().get("default_generation_settings") or {}).get("n_ctx") + child_ctx = _n_ctx(pr.json()) if child_ctx: cache[child_id]["context_length"] = child_ctx +def _remember_endpoint_models(normalized: str, cache: Dict[str, Dict[str, Any]]) -> Dict[str, Dict[str, Any]]: + _endpoint_model_metadata_cache[normalized] = cache + _endpoint_model_metadata_cache_time[normalized] = time.time() + return cache + + +def _parse_models_payload(payload: Dict[str, Any]) -> Dict[str, Dict[str, Any]]: + cache: Dict[str, Dict[str, Any]] = {} + for model in payload.get("data", []): + model_id = model.get("id") if isinstance(model, dict) else None + if model_id: + _add_model_aliases(cache, model_id, _endpoint_model_entry(model, model_id, _extract_context_length(model))) + return cache + + def fetch_endpoint_model_metadata( base_url: str, api_key: str = "", @@ -1278,45 +1235,34 @@ def fetch_endpoint_model_metadata( if not normalized or _is_openrouter_base_url(normalized): return {} _ensure_requests() + local = is_local_endpoint(normalized) if not force_refresh: cached = _endpoint_model_metadata_cache.get(normalized) - cached_at = _endpoint_model_metadata_cache_time.get(normalized, 0) - if cached is not None and (time.time() - cached_at) < _ENDPOINT_MODEL_CACHE_TTL: + if cached is not None and (time.time() - _endpoint_model_metadata_cache_time.get(normalized, 0)) < _ENDPOINT_MODEL_CACHE_TTL: return cached - if not is_local_endpoint(normalized): + if not local: memo = _endpoint_disk_cache_get(normalized) if memo is not None: - _endpoint_model_metadata_cache[normalized] = memo - _endpoint_model_metadata_cache_time[normalized] = time.time() - return memo + return _remember_endpoint_models(normalized, memo) # Blackholed: return empty WITHOUT caching so it is retried once the entry expires. if _endpoint_blackholed(normalized): return {} - candidates = [normalized] - if normalized.endswith("/v1"): - alternate = normalized[:-3].rstrip("/") - else: - alternate = normalized + "/v1" - if alternate and alternate not in candidates: - candidates.append(alternate) - + alternate = normalized[:-3].rstrip("/") if normalized.endswith("/v1") else normalized + "/v1" + candidates = [normalized] + ([alternate] if alternate and alternate != normalized else []) headers = {"Authorization": f"Bearer {api_key}"} if api_key else {} + verify = _resolve_requests_verify(normalized) last_error: Optional[Exception] = None - if is_local_endpoint(normalized): + if local: try: if detect_local_server_type(normalized, api_key=api_key) == "lm-studio": - cache = _lmstudio_native_models(normalized, headers) - _endpoint_model_metadata_cache[normalized] = cache - _endpoint_model_metadata_cache_time[normalized] = time.time() - return cache + return _remember_endpoint_models(normalized, _lmstudio_native_models(normalized, headers)) except Exception as exc: last_error = exc - if _is_connect_timeout(exc): - _note_endpoint_blackholed(normalized) + _note_if_connect_timeout(exc, normalized) for candidate in candidates: # A connect timeout condemns the host, not the path. @@ -1327,13 +1273,7 @@ def fetch_endpoint_model_metadata( url = request_candidate.rstrip("/") + "/models" response = None try: - response = requests.get( - url, - headers=headers, - timeout=(5, 10), - verify=_resolve_requests_verify(normalized), - stream=True, - ) + response = requests.get(url, headers=headers, timeout=(5, 10), verify=verify, stream=True) if response.status_code in (401, 403): logger.debug( "Model metadata probe received HTTP %s from %s; stopping candidate probing", @@ -1343,46 +1283,28 @@ def fetch_endpoint_model_metadata( break response.raise_for_status() payload = response.json() - cache: Dict[str, Dict[str, Any]] = {} - for model in payload.get("data", []): - if not isinstance(model, dict): - continue - model_id = model.get("id") - if not model_id: - continue - _add_model_aliases(cache, model_id, _endpoint_model_entry(model, model_id, _extract_context_length(model))) - + cache = _parse_models_payload(payload) if any(m.get("owned_by") == "llamacpp" for m in payload.get("data", []) if isinstance(m, dict)): try: - _apply_llamacpp_props(cache, request_candidate, headers, _resolve_requests_verify(normalized)) + _apply_llamacpp_props(cache, request_candidate, headers, verify) except Exception: pass - - _endpoint_model_metadata_cache[normalized] = cache - _endpoint_model_metadata_cache_time[normalized] = time.time() - if cache and not is_local_endpoint(normalized): + if cache and not local: _endpoint_disk_cache_put(normalized, cache) - return cache + return _remember_endpoint_models(normalized, cache) except Exception as exc: last_error = exc - if _is_connect_timeout(exc): - _note_endpoint_blackholed(normalized) + _note_if_connect_timeout(exc, normalized) finally: if response is not None: response.close() if last_error: logger.debug("Failed to fetch model metadata from %s/models: %s", normalized, last_error) - _endpoint_model_metadata_cache[normalized] = {} - _endpoint_model_metadata_cache_time[normalized] = time.time() - return {} + return _remember_endpoint_models(normalized, {}) -def _resolve_endpoint_context_length( - model: str, - base_url: str, - api_key: str = "", -) -> Optional[int]: +def _resolve_endpoint_context_length(model: str, base_url: str, api_key: str = "") -> Optional[int]: """Resolve context length from an endpoint's live ``/models`` metadata.""" endpoint_metadata = fetch_endpoint_model_metadata(base_url, api_key=api_key) matched = endpoint_metadata.get(model) @@ -1391,19 +1313,13 @@ def _resolve_endpoint_context_length( matched = next(iter(endpoint_metadata.values())) elif model: # Substring match; "" would match EVERY key and poison the window. - for key, entry in endpoint_metadata.items(): - if model in key or key in model: - matched = entry - break - if matched: - context_length = matched.get("context_length") - if isinstance(context_length, int): - return context_length - return None + matched = next((entry for key, entry in endpoint_metadata.items() if model in key or key in model), None) + context_length = matched.get("context_length") if matched else None + return context_length if isinstance(context_length, int) else None def _get_context_cache_path() -> Path: - """Return path to the persistent context length cache file.""" + """Path to the persistent context length cache file.""" from hermes_constants import get_hermes_home return get_hermes_home() / "context_length_cache.yaml" @@ -1422,25 +1338,26 @@ def _load_context_cache() -> Dict[str, int]: return {} -def _context_cache_key(model: str, base_url: str) -> str: - """Canonical ``model@base_url`` key for the persistent context cache. +def _write_context_cache(cache: Dict[str, int]) -> None: + """Atomic write (temp file + fsync + os.replace). - Trailing slashes are stripped so ``http://host/v1`` and - ``http://host/v1/`` share one entry instead of creating duplicates - that can go stale independently. + A plain truncating ``open(path, "w")`` leaves the file empty/partial if the + process is killed mid-dump, and the next _load_context_cache() swallows the + YAML error and returns {} — silently wiping EVERY cached context length. It + also exposes torn reads to a concurrent reader. Raises on failure. """ + atomic_yaml_write(_get_context_cache_path(), {"context_lengths": cache}) + + +def _context_cache_key(model: str, base_url: str) -> str: + """Canonical ``model@base_url`` key; trailing slashes stripped so ``/v1`` and ``/v1/`` share one entry.""" return f"{model}@{(base_url or '').rstrip('/')}" def save_context_length(model: str, base_url: str, length: int) -> None: - """Persist a discovered context length for a model+provider combo. - - Cache key is ``model@base_url`` so the same model name served from - different providers can have different limits. - """ - # Never persist non-positive values — a 0 or negative context length - # is always a bug and would poison the cache, causing downstream - # `get_model_context_length()` to return 0 (since `0 is not None`). + """Persist a discovered context length under ``model@base_url`` (same model, different providers, different limits).""" + # Never persist non-positive values — a 0 or negative context length is + # always a bug and would make get_model_context_length() return 0 (since `0 is not None`). if length <= 0: logger.warning( "Refusing to cache non-positive context length %s -> %s tokens", @@ -1452,15 +1369,8 @@ def save_context_length(model: str, base_url: str, length: int) -> None: if cache.get(key) == length: return # already stored cache[key] = length - path = _get_context_cache_path() try: - # Atomic write (temp file + fsync + os.replace): a plain truncating - # ``open(path, "w")`` leaves the file empty/partial if the process is - # killed mid-dump, and the next _load_context_cache() swallows the - # resulting YAML error and returns {} — silently wiping EVERY cached - # context length. It also exposes torn reads to a concurrent process - # reading between truncate and dump-complete. - atomic_yaml_write(path, {"context_lengths": cache}) + _write_context_cache(cache) logger.info("Cached context length %s -> %s tokens", key, f"{length:,}") except Exception as e: logger.debug("Failed to save context length cache: %s", e) @@ -1470,19 +1380,13 @@ def get_cached_context_length(model: str, base_url: str) -> Optional[int]: """Look up a previously discovered context length for model+provider.""" key = _context_cache_key(model, base_url) cache = _load_context_cache() - hit = cache.get(key) - if hit is not None: - return hit - # Legacy rows written before key normalization may carry a trailing - # slash — honor them rather than re-probing. Checked regardless of the - # caller's slash form: the row's shape and the caller's shape can differ - # in either direction (old slashed row + new normalized config, or the - # reverse), so probe the literal form and the slashed canonical form. - for legacy_key in (f"{model}@{base_url}", f"{key}/"): - if legacy_key != key: - hit = cache.get(legacy_key) - if hit is not None: - return hit + # Legacy rows written before key normalization may carry a trailing slash; + # the row's shape and the caller's can differ in either direction, so probe + # the canonical key, the literal form and the slashed canonical form. + for candidate in (key, f"{model}@{base_url}", f"{key}/"): + hit = cache.get(candidate) + if hit is not None: + return hit return None @@ -1490,42 +1394,34 @@ def _invalidate_cached_context_length(model: str, base_url: str) -> None: """Drop a stale cache entry so it gets re-resolved on the next lookup.""" key = _context_cache_key(model, base_url) cache = _load_context_cache() - # Invalidation must also drop the in-memory TTL probe entries for this - # pair — otherwise the next resolution inside the TTL window reuses the - # very value we just declared stale and re-persists it. + # Also drop the in-memory TTL probe entries for this pair — otherwise the + # next resolution inside the TTL window reuses the value just declared stale. bare = _strip_provider_prefix(model) stripped = (base_url or "").rstrip("/") _LOCAL_CTX_PROBE_CACHE.pop((bare, stripped), None) _LOCAL_CTX_PROBE_CACHE.pop(("ollama_show", bare, stripped), None) - # Clear every key shape for this pair: canonical, the caller's literal - # form, and the slashed legacy form — same set get_cached_context_length - # consults, so a lookup can never resurrect a row invalidation missed. + # Every key shape get_cached_context_length consults, so a lookup can never + # resurrect a row invalidation missed. stale_keys = {key, f"{model}@{base_url}", f"{key}/"} if not any(k in cache for k in stale_keys): return for k in stale_keys: cache.pop(k, None) - path = _get_context_cache_path() try: - # Atomic write — see save_context_length() for why a plain truncating - # open() here risks wiping the entire cache on an interrupted dump. - atomic_yaml_write(path, {"context_lengths": cache}) + _write_context_cache(cache) except Exception as e: logger.debug("Failed to invalidate context length cache entry %s: %s", key, e) def get_next_probe_tier(current_length: int) -> Optional[int]: - """Return the next lower probe tier, or None if already at minimum.""" - for tier in CONTEXT_PROBE_TIERS: - if tier < current_length: - return tier - return None + """Next lower probe tier, or None if already at minimum.""" + return next((tier for tier in CONTEXT_PROBE_TIERS if tier < current_length), None) def parse_context_limit_from_error(error_msg: str) -> Optional[int]: """Context limit quoted in a provider error ("maximum context length is 32768 tokens"), if any.""" error_lower = error_msg.lower() - patterns = [ + patterns = ( r'max_model_len\s*(?:is\s*)?[:=(]?\s*(\d{4,})', # vLLM: "max_model_len 32768", "=32768", ": 32768", "(32768)", "is 32768" r'maximum model length\s*(?:is\s*)?[:=(]?\s*(\d{4,})', # vLLM alt: "maximum model length 131072", "... is 131072" r'(?:max(?:imum)?|limit)\s*(?:context\s*)?(?:length|size|window)?\s*(?:is|of|:)?\s*(\d{4,})', @@ -1536,21 +1432,17 @@ def parse_context_limit_from_error(error_msg: str) -> Optional[int]: # Gemini: "input token count is 32825 but model only supports up to # 32768" — anchor on the phrase so the input count isn't captured. r'supports?\s+(?:only\s+)?up\s+to\s+(\d{4,})', - ] + ) for pattern in patterns: match = re.search(pattern, error_lower) if match: limit = int(match.group(1)) - # Sanity check: must be a reasonable context length - if 1024 <= limit <= 10_000_000: + if 1024 <= limit <= 10_000_000: # sanity: must be a plausible window return limit return None -def get_context_length_from_provider_error( - error_msg: str, - current_context_length: int, -) -> Optional[int]: +def get_context_length_from_provider_error(error_msg: str, current_context_length: int) -> Optional[int]: """Provider-reported limit LOWER than the current window, else None. Overflow recovery must not invent a window: when the provider only says @@ -1608,9 +1500,7 @@ def parse_available_output_tokens_from_error(error_msg: str) -> Optional[int]: _m_ctx_tok = re.search(r'maximum context length is (\d+)\s*token', error_lower) _m_chars = re.search(r'prompt contains (\d+)\s*character', error_lower) if _m_ctx_tok and _m_chars: - _ctx = int(_m_ctx_tok.group(1)) - _est_input = (int(_m_chars.group(1)) + 2) // 3 - _available = _ctx - _est_input + _available = int(_m_ctx_tok.group(1)) - (int(_m_chars.group(1)) + 2) // 3 if _available >= 1: return _available @@ -1621,9 +1511,7 @@ def parse_available_output_tokens_from_error(error_msg: str) -> Optional[int]: # requested_output - 1 and each retry walks the cap down by the safety # margin without ever fitting. Detect that and halve the cap instead — # still strictly below what was rejected, converges in one or two retries. - _m_vllm_input = re.search( - r'prompt contains (?:at least )?(\d+)\s*input tokens', error_lower - ) + _m_vllm_input = re.search(r'prompt contains (?:at least )?(\d+)\s*input tokens', error_lower) if _m_ctx_tok and _m_vllm_input: _available = int(_m_ctx_tok.group(1)) - int(_m_vllm_input.group(1)) _m_requested_out = re.search(r'requested (\d+)\s*output tokens', error_lower) @@ -1682,13 +1570,13 @@ def is_output_cap_error(error_msg: str) -> bool: NOT about the input being too long; when both appear, defer to overflow. """ error_lower = error_msg.lower() - if not any(p in error_lower for p in ("max_tokens", "max_output_tokens", "max_completion_tokens")): - return False - if not _any_phrase_group(error_lower, _OUTPUT_CAP_SIGNALS): - return False - # An error that ALSO describes an oversized INPUT is a genuine context - # overflow that happens to mention max_tokens — compression can fix it. - return not any(p in error_lower for p in _INPUT_OVERFLOW_SIGNALS) + return ( + any(p in error_lower for p in ("max_tokens", "max_output_tokens", "max_completion_tokens")) + and _any_phrase_group(error_lower, _OUTPUT_CAP_SIGNALS) + # An error that ALSO describes an oversized INPUT is a genuine context + # overflow that happens to mention max_tokens — compression can fix it. + and not any(p in error_lower for p in _INPUT_OVERFLOW_SIGNALS) + ) def _model_id_matches(candidate_id: str, lookup_model: str) -> bool: @@ -1698,18 +1586,30 @@ def _model_id_matches(candidate_id: str, lookup_model: str) -> bool: ) -def query_ollama_num_ctx(model: str, base_url: str, api_key: str = "") -> Optional[int]: - """Ollama ``/api/show`` context (Modelfile num_ctx, else GGUF max); the value to send as ``num_ctx``.""" +def _ollama_show(server_url: str, api_key: str, bare_model: str) -> Optional[Dict[str, Any]]: + """Ollama ``/api/show`` JSON for ``bare_model`` (3 s timeout), or None on any failure.""" import httpx - bare_model = _strip_provider_prefix(model) - server_url = _server_root(base_url) - try: - server_type = detect_local_server_type(base_url, api_key=api_key) + with httpx.Client(timeout=3.0, headers=_auth_headers(api_key)) as client: + resp = client.post(f"{server_url}/api/show", json={"name": bare_model}) + return resp.json() if resp.status_code == 200 else None except Exception: return None - if server_type != "ollama": + + +def _is_ollama_server(base_url: str, api_key: str) -> bool: + try: + return detect_local_server_type(base_url, api_key=api_key) == "ollama" + except Exception: + return False + + +def query_ollama_num_ctx(model: str, base_url: str, api_key: str = "") -> Optional[int]: + """Ollama ``/api/show`` context (Modelfile num_ctx, else GGUF max); the value to send as ``num_ctx``.""" + bare_model = _strip_provider_prefix(model) + server_url = _server_root(base_url) + if not _is_ollama_server(base_url, api_key): return None _disk_key = f"{server_url}|{bare_model}" @@ -1717,51 +1617,25 @@ def query_ollama_num_ctx(model: str, base_url: str, api_key: str = "") -> Option if isinstance(disk_hit, int) and disk_hit > 0: return disk_hit - headers = _auth_headers(api_key) - - try: - with httpx.Client(timeout=3.0, headers=headers) as client: - resp = client.post(f"{server_url}/api/show", json={"name": bare_model}) - if resp.status_code != 200: - return None - ctx = _ollama_show_context(resp.json(), gguf_first=False) - if ctx is not None: - _local_probe_disk_put("ollama_num_ctx", _disk_key, ctx) - return ctx - except Exception: - pass - return None + data = _ollama_show(server_url, api_key, bare_model) + ctx = _ollama_show_context(data, gguf_first=False) if data is not None else None + if ctx is not None: + _local_probe_disk_put("ollama_num_ctx", _disk_key, ctx) + return ctx def query_ollama_supports_vision(model: str, base_url: str, api_key: str = "") -> Optional[bool]: - """Return True/False when Ollama ``/api/show`` reports vision support. + """True/False when Ollama ``/api/show`` reports vision support, None when unknown. Uses the ``capabilities`` field on Ollama 0.6.0+ and falls back to - ``model_info.*.vision.block_count`` on older servers. Returns None when - the server is unreachable, not Ollama, or the model is unknown. + ``model_info.*.vision.block_count`` on older servers. None when the server + is unreachable, not Ollama, or the model is unknown. """ - import httpx - bare_model = _strip_provider_prefix(model) - if not bare_model or not base_url: + if not bare_model or not base_url or not _is_ollama_server(base_url, api_key): return None - - try: - if detect_local_server_type(base_url, api_key=api_key) != "ollama": - return None - except Exception: - return None - - server_url = _server_root(base_url) - headers = _auth_headers(api_key) - - try: - with httpx.Client(timeout=3.0, headers=headers) as client: - resp = client.post(f"{server_url}/api/show", json={"name": bare_model}) - if resp.status_code != 200: - return None - data = resp.json() - except Exception: + data = _ollama_show(_server_root(base_url), api_key, bare_model) + if data is None: return None caps = data.get("capabilities") @@ -1772,14 +1646,27 @@ def query_ollama_supports_vision(model: str, base_url: str, api_key: str = "") - return False model_info = data.get("model_info") - if isinstance(model_info, dict): - for key in model_info: - if "vision.block_count" in str(key).lower(): - return True - + if isinstance(model_info, dict) and any("vision.block_count" in str(key).lower() for key in model_info): + return True return None +def _memo_local_probe(cache_key: tuple, probe: Callable[[], Optional[int]]) -> Optional[int]: + """Short-TTL, positive-only memo of a local probe (see _LOCAL_CTX_PROBE_CACHE). + + A failure during a startup race must not suppress the retry seconds later + once the server is up, so only truthy results are memoized. + """ + now = time.monotonic() + cached = _LOCAL_CTX_PROBE_CACHE.get(cache_key) + if cached is not None and (now - cached[1]) < _LOCAL_CTX_PROBE_TTL_SECONDS: + return cached[0] + result = probe() + if result: + _LOCAL_CTX_PROBE_CACHE[cache_key] = (result, now) + return result + + def _query_ollama_api_show(model: str, base_url: str, api_key: str = "") -> Optional[int]: """Provider-agnostic Ollama ``/api/show`` context probe (any hostname; non-Ollama servers 404 fast). @@ -1787,18 +1674,10 @@ def _query_ollama_api_show(model: str, base_url: str, api_key: str = "") -> Opti query_ollama_num_ctx(). Positive results share _LOCAL_CTX_PROBE_CACHE under a namespaced key (the two probes can differ for the same (model, url)). """ - import time as _time - - cache_key = ("ollama_show", _strip_provider_prefix(model), base_url.rstrip("/")) - now = _time.monotonic() - cached = _LOCAL_CTX_PROBE_CACHE.get(cache_key) - if cached is not None and (now - cached[1]) < _LOCAL_CTX_PROBE_TTL_SECONDS: - return cached[0] - - result = _query_ollama_api_show_uncached(model, base_url, api_key=api_key) - if result: # positive-only — never memoize a failed probe - _LOCAL_CTX_PROBE_CACHE[cache_key] = (result, now) - return result + return _memo_local_probe( + ("ollama_show", _strip_provider_prefix(model), base_url.rstrip("/")), + lambda: _query_ollama_api_show_uncached(model, base_url, api_key=api_key), + ) def _query_ollama_api_show_uncached(model: str, base_url: str, api_key: str = "") -> Optional[int]: @@ -1808,22 +1687,16 @@ def _query_ollama_api_show_uncached(model: str, base_url: str, api_key: str = "" server_url = _server_root(base_url) if _endpoint_blackholed(server_url): return None - - headers = _auth_headers(api_key) - try: - with httpx.Client(timeout=5.0, headers=headers) as client: + with httpx.Client(timeout=5.0, headers=_auth_headers(api_key)) as client: resp = client.post(f"{server_url}/api/show", json={"name": model}) if resp.status_code != 200: return None # Hosted Ollama: the GGUF max is authoritative (the operator may # have capped num_ctx arbitrarily). - ctx = _ollama_show_context(resp.json(), gguf_first=True, minimum=1024) - if ctx is not None: - return ctx + return _ollama_show_context(resp.json(), gguf_first=True, minimum=1024) except Exception as exc: - if _is_connect_timeout(exc): - _note_endpoint_blackholed(server_url) + _note_if_connect_timeout(exc, server_url) return None @@ -1862,21 +1735,14 @@ def _stale_pre_catalog_cache_entry(model: str, cached: int) -> bool: 256K fallback). Values above that — genuine probe results — are kept. """ model_lower = model.lower() - matches = [ - (key, value) - for key, value in DEFAULT_CONTEXT_LENGTHS.items() - if key in model_lower - ] + matches = [(key, value) for key, value in DEFAULT_CONTEXT_LENGTHS.items() if key in model_lower] if not matches: return False specific_key, specific_value = max(matches, key=lambda kv: len(kv[0])) - if specific_key not in _PRE_CATALOG_STALE_KEYS: - return False - if cached >= specific_value: + if specific_key not in _PRE_CATALOG_STALE_KEYS or cached >= specific_value: return False shorter_values = [v for k, v in matches if len(k) < len(specific_key)] - threshold = max(shorter_values, default=DEFAULT_FALLBACK_CONTEXT) - return cached <= threshold + return cached <= max(shorter_values, default=DEFAULT_FALLBACK_CONTEXT) def _model_name_suggests_minimax(model: str) -> bool: @@ -1886,26 +1752,81 @@ def _model_name_suggests_minimax(model: str) -> bool: def _model_name_suggests_stale_32k_underreport(model: str) -> bool: - """Return True for model families known to be wrongly underreported as 32K.""" + """Model families known to be wrongly underreported as 32K.""" return _model_name_suggests_kimi(model) or _model_name_suggests_minimax(model) def _query_local_context_length(model: str, base_url: str, api_key: str = "") -> Optional[int]: """Local-server context probe, short-TTL cached (see _LOCAL_CTX_PROBE_CACHE).""" - import time as _time + return _memo_local_probe( + (_strip_provider_prefix(model), base_url.rstrip("/")), + lambda: _query_local_context_length_uncached(model, base_url, api_key=api_key), + ) - cache_key = (_strip_provider_prefix(model), base_url.rstrip("/")) - now = _time.monotonic() - cached = _LOCAL_CTX_PROBE_CACHE.get(cache_key) - if cached is not None and (now - cached[1]) < _LOCAL_CTX_PROBE_TTL_SECONDS: - return cached[0] - result = _query_local_context_length_uncached(model, base_url, api_key=api_key) - # Positive-only: a failure during a startup race must not suppress the - # retry seconds later once the server is up. - if result: - _LOCAL_CTX_PROBE_CACHE[cache_key] = (result, now) - return result +def _positive_int(value: Any) -> Optional[int]: + return int(value) if isinstance(value, (int, float)) and value else None + + +def _lmstudio_context(client, lmstudio_url: str, model: str) -> Optional[int]: + """LM Studio native /api/v1/models (the OpenAI-compat list omits context); + loaded-instance config is the runtime value.""" + resp = client.get(f"{lmstudio_url}/api/v1/models") + if resp.status_code != 200: + return None + for m in resp.json().get("models", []): + if _model_id_matches(m.get("key", ""), model) or _model_id_matches(m.get("id", ""), model): + for inst in m.get("loaded_instances", []): + ctx = _positive_int(inst.get("config", {}).get("context_length")) + if ctx is not None: + return ctx + return None + return None + + +def _llamacpp_context(client, server_url: str, model: str) -> Optional[int]: + """llama.cpp /props: the RUNTIME n_ctx, answered by the router even for a + not-yet-loaded model (while /v1/models has meta=null), so a lazily-loaded + model doesn't fall to a family catch-all.""" + import httpx + + for props_path in (f"/props?model={model}", "/props"): + try: + resp = client.get(f"{server_url}{props_path}") + except httpx.HTTPError: + return None + if resp.status_code != 200: + continue + n_ctx = _positive_int((resp.json().get("default_generation_settings") or {}).get("n_ctx")) + if n_ctx is not None: + return n_ctx + return None + + +def _openai_models_list_context(client, server_url: str, model: str) -> Optional[int]: + """/v1/models list: match by id, else the sole model on single-model servers + (llama.cpp reports a GGUF path as id, which rarely equals the configured name).""" + resp = client.get(f"{server_url}/v1/models") + if resp.status_code != 200: + return None + models_list = resp.json().get("data", []) + matched = next((m for m in models_list if isinstance(m, dict) and _model_id_matches(m.get("id", ""), model)), None) + if matched is None and len(models_list) == 1: + matched = models_list[0] + if matched is None: + return None + # Runtime n_ctx (llama.cpp nests it under meta) beats n_ctx_train, which + # can exceed what the server allocates. + sources = [s for s in (matched, matched.get("meta") or {}) if isinstance(s, dict)] + for source in sources: + val = _positive_int(source.get("n_ctx")) + if val is not None: + return val + for source in sources: + ctx = _context_length_from_model_payload(source) + if ctx is not None: + return ctx + return None def _query_local_context_length_uncached(model: str, base_url: str, api_key: str = "") -> Optional[int]: @@ -1913,22 +1834,18 @@ def _query_local_context_length_uncached(model: str, base_url: str, api_key: str import httpx model = _strip_provider_prefix(model) - server_url = _server_root(base_url) lmstudio_url = _localhost_to_ipv4(_lmstudio_server_root(base_url)) - if _endpoint_blackholed(server_url): return None - headers = _auth_headers(api_key) - try: server_type = detect_local_server_type(base_url, api_key=api_key) except Exception: server_type = None try: - with httpx.Client(timeout=3.0, headers=headers) as client: + with httpx.Client(timeout=3.0, headers=_auth_headers(api_key)) as client: # Ollama: num_ctx (runtime window) before the GGUF training max — # using the max would let conversations grow past what Ollama # allocated and it would silently truncate. Matches query_ollama_num_ctx(). @@ -1938,87 +1855,25 @@ def _query_local_context_length_uncached(model: str, base_url: str, api_key: str ctx = _ollama_show_context(resp.json(), gguf_first=False) if ctx is not None: return ctx + elif server_type == "lm-studio": + ctx = _lmstudio_context(client, lmstudio_url, model) + if ctx is not None: + return ctx + elif server_type == "llamacpp": + ctx = _llamacpp_context(client, server_url, model) + if ctx is not None: + return ctx - # LM Studio native /api/v1/models (the OpenAI-compat list omits - # context); loaded-instance config is the runtime value. - if server_type == "lm-studio": - resp = client.get(f"{lmstudio_url}/api/v1/models") - if resp.status_code == 200: - data = resp.json() - for m in data.get("models", []): - if _model_id_matches(m.get("key", ""), model) or _model_id_matches(m.get("id", ""), model): - # Prefer loaded instance context (actual runtime value) - for inst in m.get("loaded_instances", []): - cfg = inst.get("config", {}) - ctx = cfg.get("context_length") - if ctx and isinstance(ctx, (int, float)): - return int(ctx) - break - - # llama.cpp /props: the RUNTIME n_ctx, answered by the router even - # for a not-yet-loaded model (while /v1/models has meta=null), so a - # lazily-loaded model doesn't fall to a family catch-all. - if server_type == "llamacpp": - for props_path in (f"/props?model={model}", "/props"): - try: - resp = client.get(f"{server_url}{props_path}") - except httpx.HTTPError: - break - if resp.status_code != 200: - continue - n_ctx = (resp.json().get("default_generation_settings") - or {}).get("n_ctx") - if isinstance(n_ctx, (int, float)) and n_ctx: - return int(n_ctx) - - # LM Studio / vLLM / llama.cpp / Anthropic-compat proxies: - # try /v1/models/{model} + # LM Studio / vLLM / llama.cpp / Anthropic-compat proxies: /v1/models/{model} resp = client.get(f"{server_url}/v1/models/{model}") if resp.status_code == 200: - data = resp.json() - if isinstance(data, dict): - ctx = _context_length_from_model_payload(data) - if ctx is not None: - return ctx + ctx = _context_length_from_model_payload(resp.json()) + if ctx is not None: + return ctx - # Try /v1/models and find the model in the list. - # Use _model_id_matches to handle "publisher/slug" vs bare "slug". - resp = client.get(f"{server_url}/v1/models") - if resp.status_code == 200: - data = resp.json() - models_list = data.get("data", []) - # Match by id; on single-model servers (e.g. llama.cpp) the - # configured name rarely equals the reported id (a GGUF path), - # so fall back to the sole model when nothing matches. - matched = None - for m in models_list: - if not isinstance(m, dict): - continue - if _model_id_matches(m.get("id", ""), model): - matched = m - break - if matched is None and len(models_list) == 1: - matched = models_list[0] - if matched is not None: - # Runtime n_ctx (llama.cpp nests it under meta) beats - # n_ctx_train, which can exceed what the server allocates. - sources = [ - s - for s in (matched, matched.get("meta") or {}) - if isinstance(s, dict) - ] - for source in sources: - val = source.get("n_ctx") - if isinstance(val, (int, float)) and val: - return int(val) - for source in sources: - ctx = _context_length_from_model_payload(source) - if ctx is not None: - return ctx + return _openai_models_list_context(client, server_url, model) except Exception as exc: - if _is_connect_timeout(exc): - _note_endpoint_blackholed(server_url) - + _note_if_connect_timeout(exc, server_url) return None @@ -2035,17 +1890,12 @@ def _query_anthropic_context_length(model: str, base_url: str, api_key: str) -> base = base_url.rstrip("/") if base.endswith("/v1"): base = base[:-3] - url = f"{base}/v1/models?limit=1000" - headers = { - "x-api-key": api_key, - "anthropic-version": "2023-06-01", - } + headers = {"x-api-key": api_key, "anthropic-version": "2023-06-01"} _ensure_requests() - resp = requests.get(url, headers=headers, timeout=(5, 10), verify=_resolve_requests_verify(base_url)) + resp = requests.get(f"{base}/v1/models?limit=1000", headers=headers, timeout=(5, 10), verify=_resolve_requests_verify(base_url)) if resp.status_code != 200: return None - data = resp.json() - for m in data.get("data", []): + for m in resp.json().get("data", []): if m.get("id") == model: ctx = m.get("max_input_tokens") if isinstance(ctx, int) and ctx > 0: @@ -2076,9 +1926,8 @@ _CODEX_OAUTH_CONTEXT_FALLBACK: Dict[str, int] = { } # Codex OAuth advertises 272K for these families but ACCEPTS ~900K+ (verified -# live: 911,276 input tokens OK on gpt-5.6-sol; terra/luna/gpt-5.4 completed -# 900,026; gpt-5.5 and gpt-5.4-mini genuinely reject >272K and are NOT -# listed). 900K keeps ≥11K margin under the observed ceiling. +# live; gpt-5.5 and gpt-5.4-mini genuinely reject >272K and are NOT listed). +# 900K keeps ≥11K margin under the observed ceiling. # # OPT-IN ONLY: the large window is exposed via explicit ``-900k`` picker # variants; base slugs keep 272K so the cheaper limit is the default (a 900K @@ -2086,10 +1935,9 @@ _CODEX_OAUTH_CONTEXT_FALLBACK: Dict[str, int] = { # a Hermes-side alias stripped before the wire (strip_codex_context_variant_suffix). # # The bump fires ONLY when the resolved value is exactly the stale 272,000 -# advertisement; any other advertised number is trusted and the table is -# inert. ``gpt-5.6`` is a FAMILY PREFIX (``-pro`` slugs aren't routable on -# Codex, so over-matching is moot); ``gpt-5.4`` is EXACT because gpt-5.4-mini -# enforces 272K. +# advertisement; any other advertised number is trusted. ``gpt-5.6`` is a +# FAMILY PREFIX (``-pro`` slugs aren't routable on Codex, so over-matching is +# moot); ``gpt-5.4`` is EXACT because gpt-5.4-mini enforces 272K. _CODEX_OAUTH_VERIFIED_ABOVE_ADVERTISED_PREFIXES: Dict[str, int] = { "gpt-5.6": 900_000, # sol / terra / luna } @@ -2129,20 +1977,16 @@ def is_codex_900k_base(model: Optional[str]) -> bool: if slug in _CODEX_900K_ELIGIBLE_BASES: return True # Dated snapshots of the routable 5.6 bases (gpt-5.6-sol-2026-07-09). - for base in _CODEX_900K_SNAPSHOT_BASES: - if slug.startswith(base + "-") and _CODEX_900K_SNAPSHOT_RE.match( - slug[len(base) + 1:] - ): - return True - return False + return any( + slug.startswith(base + "-") and _CODEX_900K_SNAPSHOT_RE.match(slug[len(base) + 1:]) + for base in _CODEX_900K_SNAPSHOT_BASES + ) def is_codex_context_variant(model: Optional[str]) -> bool: """Suffix AND eligible base — ``gpt-5.5-900k`` is an invalid alias, not a variant.""" slug = _bare_codex_slug(model) - if not slug.endswith(CODEX_CONTEXT_VARIANT_SUFFIX): - return False - return is_codex_900k_base(slug[: -len(CODEX_CONTEXT_VARIANT_SUFFIX)]) + return slug.endswith(CODEX_CONTEXT_VARIANT_SUFFIX) and is_codex_900k_base(slug[: -len(CODEX_CONTEXT_VARIANT_SUFFIX)]) def strip_codex_context_variant_suffix(model: Optional[str]) -> str: @@ -2152,11 +1996,10 @@ def strip_codex_context_variant_suffix(model: Optional[str]) -> str: honestly at the API instead of silently running as a different model. """ raw = (model or "").strip() - if not raw.lower().endswith(CODEX_CONTEXT_VARIANT_SUFFIX): - return raw - base = raw[: -len(CODEX_CONTEXT_VARIANT_SUFFIX)] - if is_codex_900k_base(base): - return base + if raw.lower().endswith(CODEX_CONTEXT_VARIANT_SUFFIX): + base = raw[: -len(CODEX_CONTEXT_VARIANT_SUFFIX)] + if is_codex_900k_base(base): + return base return raw @@ -2187,7 +2030,7 @@ _CODEX_OAUTH_CONTEXT_CACHE_TTL = 3600 # 1 hour def _codex_oauth_token_fingerprint(access_token: str) -> str: - """Return a non-secret cache key for a Codex OAuth access token.""" + """Non-secret cache key for a Codex OAuth access token.""" return hashlib.sha256(access_token.encode("utf-8")).hexdigest()[:16] @@ -2211,23 +2054,18 @@ def _extract_chatgpt_account_id(access_token: str) -> Optional[str]: return None -def _fetch_codex_oauth_context_lengths_with_source( - access_token: str, -) -> Tuple[Dict[str, int], bool]: +def _fetch_codex_oauth_context_lengths_with_source(access_token: str) -> Tuple[Dict[str, int], bool]: """Codex catalogue ``{slug: context_window}`` plus whether it came from HTTP. Cached per token fingerprint (windows vary by entitlement; the raw token is never a key). An in-process hit reports False: it is not a fresh provider confirmation and must not drive persistent writes. """ - global _codex_oauth_context_cache now = time.time() cache_key = _codex_oauth_token_fingerprint(access_token) cached = _codex_oauth_context_cache.get(cache_key) - if cached is not None: - cached_models, cached_at = cached - if now - cached_at < _CODEX_OAUTH_CONTEXT_CACHE_TTL: - return cached_models, False + if cached is not None and now - cached[1] < _CODEX_OAUTH_CONTEXT_CACHE_TTL: + return cached[0], False headers = {"Authorization": f"Bearer {access_token}"} acct_id = _extract_chatgpt_account_id(access_token) @@ -2253,9 +2091,8 @@ def _fetch_codex_oauth_context_lengths_with_source( logger.debug("Codex /models probe failed: %s", exc) return {}, False - entries = data.get("models", []) if isinstance(data, dict) else [] result: Dict[str, int] = {} - for item in entries: + for item in data.get("models", []) if isinstance(data, dict) else []: if not isinstance(item, dict): continue slug = item.get("slug") @@ -2268,9 +2105,7 @@ def _fetch_codex_oauth_context_lengths_with_source( return result, True -def _resolve_codex_oauth_context_length_with_source( - model: str, access_token: str = "" -) -> Tuple[Optional[int], str]: +def _resolve_codex_oauth_context_length_with_source(model: str, access_token: str = "") -> Tuple[Optional[int], str]: """``(context_length, source)`` for a Codex OAuth slug. source: "live" (fresh authenticated probe — the only one eligible for @@ -2301,9 +2136,8 @@ def _resolve_codex_oauth_context_length_with_source( if lookup_bare in live: return _apply_verified_bump(live[lookup_bare], live_source) # Case-insensitive match in case casing drifts - model_lower = lookup_bare.lower() for slug, ctx in live.items(): - if slug.lower() == model_lower: + if slug.lower() == lookup_bare.lower(): return _apply_verified_bump(ctx, live_source) hit = _longest_key_match(_CODEX_OAUTH_CONTEXT_FALLBACK, lookup_bare.lower()) @@ -2312,11 +2146,7 @@ def _resolve_codex_oauth_context_length_with_source( return None, "" -def _resolve_nous_context_length( - model: str, - base_url: str = "", - api_key: str = "", -) -> Tuple[Optional[int], str]: +def _resolve_nous_context_length(model: str, base_url: str = "", api_key: str = "") -> Tuple[Optional[int], str]: """``(context_length, source)`` for a Nous Portal model. Portal /v1/models is authoritative ("portal") and may differ from OR (OR @@ -2351,22 +2181,25 @@ def _resolve_nous_context_length( if ctx is not None: return ctx, "openrouter" + model_lower = model.lower() normalized = _normalize_model_version(model).lower() + def _bare(or_id: str) -> str: + return or_id.split("/", 1)[1] if "/" in or_id else or_id + + # Exact bare-id match (with dot/dash normalisation), then prefix match on a + # separator boundary — two passes so any exact hit beats every prefix hit. for or_id, entry in metadata.items(): - bare = or_id.split("/", 1)[1] if "/" in or_id else or_id - if bare.lower() == model.lower() or _normalize_model_version(bare).lower() == normalized: + bare = _bare(or_id) + if bare.lower() == model_lower or _normalize_model_version(bare).lower() == normalized: ctx = _safe_ctx(or_id, entry) if ctx is not None: return ctx, "openrouter" - model_lower = model.lower() for or_id, entry in metadata.items(): - bare = or_id.split("/", 1)[1] if "/" in or_id else or_id - for candidate, query in [(bare.lower(), model_lower), (_normalize_model_version(bare).lower(), normalized)]: - if candidate.startswith(query) and ( - len(candidate) == len(query) or candidate[len(query)] in "-:." - ): + bare = _bare(or_id) + for candidate, query in ((bare.lower(), model_lower), (_normalize_model_version(bare).lower(), normalized)): + if candidate.startswith(query) and (len(candidate) == len(query) or candidate[len(query)] in "-:."): ctx = _safe_ctx(or_id, entry) if ctx is not None: return ctx, "openrouter" @@ -2457,10 +2290,7 @@ def _resolve_bedrock_context_length(model: str, base_url: str) -> Optional[int]: base_url, else a synthetic bedrock:// key so display/offline paths share it. """ try: - from agent.bedrock_adapter import ( - get_bedrock_context_length, - resolve_bedrock_region, - ) + from agent.bedrock_adapter import get_bedrock_context_length, resolve_bedrock_region except ImportError: return None # boto3 not installed — fall through to generic resolution cache_key_url = base_url or "bedrock://" @@ -2469,11 +2299,8 @@ def _resolve_bedrock_context_length(model: str, base_url: str) -> Optional[int]: return cached # Region from the base_url host first, then the standard AWS chain. An # empty region disables probing (table only). - region = "" - if base_url: - _m = re.search(r"bedrock-runtime\.([a-z0-9-]+)\.", base_url) - if _m: - region = _m.group(1) + _m = re.search(r"bedrock-runtime\.([a-z0-9-]+)\.", base_url) if base_url else None + region = _m.group(1) if _m else "" if not region: try: region = resolve_bedrock_region() @@ -2502,8 +2329,7 @@ def _resolve_custom_endpoint_context_length(model: str, base_url: str, api_key: # 2b. Ollama native /api/show (GGUF-first for non-local). Non-Ollama servers 404/405 quickly. ctx = _query_ollama_api_show(model, base_url, api_key=api_key) if ctx is not None: - if not _skip_persistent_context_cache(base_url, provider): - save_context_length(model, base_url, ctx) + _save_unless_skipped(model, base_url, ctx, provider) return ctx # 3. Probe-down fallback after endpoint-specific detection failed logger.info( @@ -2529,81 +2355,48 @@ def _resolve_custom_endpoint_context_length(model: str, base_url: str, api_key: return DEFAULT_FALLBACK_CONTEXT -def get_model_context_length( - model: str, - base_url: str = "", - api_key: str = "", - config_context_length: int | None = None, - provider: str = "", - custom_providers: list | None = None, -) -> int: - """Get the context length for a model. +def _resolve_moa_context_length(model: str, custom_providers: list | None) -> Optional[int]: + """Step 0a: MoA virtual provider — ``model`` is a preset name and ``base_url`` + the local virtual endpoint, so every probe would miss. The aggregator is the + acting model — resolve its real provider+model (references are advisory and + never bound the acting context). None on any failure.""" + try: + from hermes_cli.config import get_compatible_custom_providers, load_config + from hermes_cli.moa_config import resolve_moa_preset + from hermes_cli.runtime_provider import resolve_runtime_provider - Resolution order: - 0. Explicit config override (model.context_length or custom_providers per-model) - 0b. model_overrides config (per-provider+model context_window override) - 0c. Endpoint-scoped metadata for models validated on one multiplexed endpoint - 1. Persistent cache (previously discovered via probing). Nous URLs, - LM Studio, and Codex OAuth bypass the cache here so their provider - metadata can be reconciled against the authoritative live source. - 1b. AWS Bedrock static table (must precede custom-endpoint probe) - 2. Active endpoint metadata (/models for explicit custom endpoints) - 3. Local server query (for local endpoints) - 4. Anthropic /v1/models API (API-key users only, not OAuth) - 5. Provider-aware lookups (before generic OpenRouter cache): - a. Copilot live /models API - b. Nous: live /v1/models probe first (authoritative), then OR - cache fallback with suffix/version normalisation. Only - portal-derived values are persisted to disk. - c. Codex OAuth /models probe - d. GMI /models endpoint - e. Ollama native /api/show probe (any base_url, provider-agnostic) - f. models.dev registry lookup (with :cloud/-cloud suffix fallback) - 6. OpenRouter live API metadata (Kimi-family 32k guard) - 7. Local server query (before hardcoded defaults for local endpoints) - 8. Hardcoded defaults (broad family patterns, longest-key-first) - 9. Default fallback (256K)""" - # 0. Explicit config override — user knows best - if config_context_length is not None and isinstance(config_context_length, int) and config_context_length > 0: - return config_context_length - - # 0a. MoA virtual provider: ``model`` is a preset name and ``base_url`` the - # local virtual endpoint, so every probe would miss. The aggregator is the - # acting model — resolve its real provider+model (references are advisory - # and never bound the acting context). Falls through on failure. - if (provider or "").strip().lower() == "moa": - try: - from hermes_cli.config import ( - get_compatible_custom_providers, - load_config, + config = load_config() + if custom_providers is None: + custom_providers = get_compatible_custom_providers(config) + agg = resolve_moa_preset(config.get("moa") or {}, model).get("aggregator") or {} + agg_provider = str(agg.get("provider") or "").strip() + agg_model = str(agg.get("model") or "").strip() + if agg_model and agg_provider and agg_provider.lower() != "moa": + rt = resolve_runtime_provider(requested=agg_provider, target_model=agg_model) + return get_model_context_length( + agg_model, + base_url=rt.get("base_url", "") or "", + api_key=rt.get("api_key", "") or "", + provider=rt.get("provider") or agg_provider, + custom_providers=custom_providers, ) - from hermes_cli.moa_config import resolve_moa_preset - from hermes_cli.runtime_provider import resolve_runtime_provider + except Exception: + logger.debug("MoA aggregator context-length resolution failed", exc_info=True) + return None - config = load_config() - effective_custom_providers = custom_providers - if effective_custom_providers is None: - effective_custom_providers = get_compatible_custom_providers(config) - preset = resolve_moa_preset(config.get("moa") or {}, model) - agg = preset.get("aggregator") or {} - agg_provider = str(agg.get("provider") or "").strip() - agg_model = str(agg.get("model") or "").strip() - if agg_model and agg_provider and agg_provider.lower() != "moa": - rt = resolve_runtime_provider(requested=agg_provider, target_model=agg_model) - return get_model_context_length( - agg_model, - base_url=rt.get("base_url", "") or "", - api_key=rt.get("api_key", "") or "", - provider=rt.get("provider") or agg_provider, - custom_providers=effective_custom_providers, - ) - except Exception: - logger.debug("MoA aggregator context-length resolution failed", exc_info=True) - # 0b. model_overrides: EXPLICIT per-provider+model context_window only. - # Fill-gap _default entries apply later inside lookup_models_dev_context - # (step 5f) once the catalog has missed, so a _default can never preempt - # custom_providers or live probes. Config-read only; never touches the network. +def _config_override_context_length( + model: str, base_url: str, provider: str, custom_providers: list | None, +) -> Optional[int]: + """Steps 0b-0c: config-only overrides (never touch the network). + + 0b. model_overrides: EXPLICIT per-provider+model context_window only. + Fill-gap _default entries apply later inside lookup_models_dev_context + (step 5f) once the catalog has missed, so a _default can never preempt + custom_providers or live probes. + 0c. custom_providers per-model override — before any probe, so /model + switch and display paths honour a per-model context_length. + """ if provider and model: try: from agent.models_dev import _override_context_window @@ -2612,9 +2405,6 @@ def get_model_context_length( return mo_ctx except Exception: pass # fall through to other resolution paths - - # 0c. custom_providers per-model override — before any probe, so /model - # switch and display paths honour a per-model context_length. if custom_providers and base_url and model: try: from hermes_cli.config import get_custom_provider_context_length @@ -2627,14 +2417,133 @@ def get_model_context_length( return cp_ctx except Exception: pass # fall through to probing + return None + + +def _resolve_provider_aware_context_length( + model: str, base_url: str, api_key: str, provider: str, effective_provider: str, +) -> Optional[int]: + """Step 5: provider-specific sources, tried in order; None when all miss.""" + # 5a. Copilot live /models — account-specific models (claude-opus-4.6-1m) + # absent from models.dev, and the provider-enforced limit for the rest. + if effective_provider in {"copilot", "copilot-acp", "github-copilot"}: + try: + from hermes_cli.models import get_copilot_model_context + ctx = get_copilot_model_context(model, api_key=api_key) + if ctx: + return ctx + except Exception: + pass # Fall through to models.dev + + if effective_provider == "nous": + ctx, source = _resolve_nous_context_length(model, base_url=base_url or "", api_key=api_key or "") + if ctx: + # Persist ONLY portal-derived values: an OR-fallback value cached + # on a portal blip would be frozen in by step 1 forever. + if base_url and source == "portal": + save_context_length(model, base_url, ctx) + return ctx + if effective_provider == "openai-codex": + # Codex OAuth enforces lower limits than the direct API for the same + # slug (gpt-5.5: 1.05M vs 272K); its own /models is authoritative. + codex_ctx, codex_source = _resolve_codex_oauth_context_length_with_source(model, access_token=api_key or "") + if codex_ctx: + # Only a fresh authenticated catalogue response may be persisted; + # the static fallback must not poison future probes. + if base_url and codex_source == "live": + save_context_length(model, base_url, codex_ctx) + return codex_ctx + if effective_provider == "gmi" and base_url: + # GMI exposes authoritative context_length via /models, but it is not + # in models.dev yet. Preserve that higher-fidelity endpoint lookup. + ctx = _resolve_endpoint_context_length(model, base_url, api_key=api_key) + if ctx is not None: + return ctx + # 5e. Ollama native /api/show for any base_url that is not a known + # non-Ollama provider (OpenAI-compat /v1/models omits context_length; the + # GGUF model_info is authoritative). Known hosted providers are skipped: + # the POST always 404s and cost ~300ms on the first-turn critical path. + if base_url: + inferred = _infer_provider_from_url(base_url) + if inferred is None or "ollama" in inferred: + ctx = _query_ollama_api_show(model, base_url, api_key=api_key) + if ctx is not None: + _save_unless_skipped(model, base_url, ctx, provider) + return ctx + # 5f. OpenRouter live /models — authoritative for OR-routed models and + # refreshed as new slugs ship, so it must win over models.dev (5g) and the + # family catch-all (8): otherwise a brand-new slug (claude-fable-5, 1M) + # falls through to the generic "claude": 200K entry. + if effective_provider == "openrouter": + entry = fetch_model_metadata().get(model) + if entry: + or_ctx = entry.get("context_length") + # Guard against the known OpenRouter Kimi-family 32k underreport + # (same class the hardcoded overrides exist to mitigate). + if isinstance(or_ctx, int) and or_ctx > 0 and not (or_ctx == 32768 and _model_name_suggests_kimi(model)): + return or_ctx + + if effective_provider: + from agent.models_dev import lookup_models_dev_context + ctx = lookup_models_dev_context(effective_provider, model) + if ctx: + # MiniMax M3: models.dev reports 512K but actual context is 1M. + # Prefer hardcoded catalog over stale probe value. + if _model_name_suggests_minimax_m3(model): + catalog = DEFAULT_CONTEXT_LENGTHS.get("minimax-m3") + if catalog and ctx < catalog: + logger.info( + "Rejecting models.dev context=%s for %r " + "(MiniMax-M3 underreport); using hardcoded default %s", + ctx, model, f"{catalog:,}", + ) + ctx = catalog + return ctx + return None + + +def get_model_context_length( + model: str, + base_url: str = "", + api_key: str = "", + config_context_length: int | None = None, + provider: str = "", + custom_providers: list | None = None, +) -> int: + """Context length for a model. + + Resolution order: + 0. Explicit config override; 0a MoA aggregator; 0b model_overrides; + 0c custom_providers per-model; endpoint-scoped metadata. + 1. Persistent cache (Nous, LM Studio and Codex OAuth bypass it so their + live source stays authoritative). 1b. AWS Bedrock static table. + 2-3. Custom endpoints: /models, local server probe, Ollama /api/show. + 4. Anthropic /v1/models (API-key users only, not OAuth). + 5. Provider-aware: Copilot, Nous (portal, then OR fallback; only portal + values persist), Codex OAuth, GMI, Ollama /api/show, OpenRouter live, + models.dev registry. + 6. OpenRouter metadata for unknown providers (32k underreport guard). + 7. Local server query before hardcoded defaults. 8. Hardcoded defaults + (longest-key-first). 9. Default fallback (256K).""" + # 0. Explicit config override — user knows best + if config_context_length is not None and isinstance(config_context_length, int) and config_context_length > 0: + return config_context_length + + if (provider or "").strip().lower() == "moa": + ctx = _resolve_moa_context_length(model, custom_providers) + if ctx is not None: + return ctx + + ctx = _config_override_context_length(model, base_url, provider, custom_providers) + if ctx is not None: + return ctx # Malformed URLs (e.g. unmatched IPv6 bracket) make urllib.parse raise; # treat them as an unknown endpoint so the inference layer reports the # configuration error itself. if base_url: try: - parsed_base_url = urlparse(_normalize_base_url(base_url)) - _ = parsed_base_url.port + _ = urlparse(_normalize_base_url(base_url)).port except ValueError: base_url = "" @@ -2669,9 +2578,7 @@ def get_model_context_length( if base_url and not _skip_persistent_context_cache(base_url, provider): cached = get_cached_context_length(model, base_url) if cached is not None: - validated = _validate_cached_context_length( - model, base_url, cached, is_bedrock_context, api_key=api_key, - ) + validated = _validate_cached_context_length(model, base_url, cached, is_bedrock_context, api_key=api_key) if validated is not None: return validated @@ -2698,9 +2605,7 @@ def get_model_context_length( return _resolve_custom_endpoint_context_length(model, base_url, api_key, provider) # 4. Anthropic /v1/models API (only for regular API keys, not OAuth) - if provider == "anthropic" or ( - base_url and base_url_hostname(base_url) == "api.anthropic.com" - ): + if provider == "anthropic" or (base_url and base_url_hostname(base_url) == "api.anthropic.com"): ctx = _query_anthropic_context_length(model, base_url or "https://api.anthropic.com", api_key) if ctx: return ctx @@ -2711,94 +2616,9 @@ def get_model_context_length( effective_provider = provider if base_url and (not effective_provider or effective_provider in {"openrouter", "custom"}): effective_provider = _infer_provider_from_url(base_url) or effective_provider - - # 5a. Copilot live /models — account-specific models (claude-opus-4.6-1m) - # absent from models.dev, and the provider-enforced limit for the rest. - if effective_provider in {"copilot", "copilot-acp", "github-copilot"}: - try: - from hermes_cli.models import get_copilot_model_context - ctx = get_copilot_model_context(model, api_key=api_key) - if ctx: - return ctx - except Exception: - pass # Fall through to models.dev - - if effective_provider == "nous": - ctx, source = _resolve_nous_context_length( - model, base_url=base_url or "", api_key=api_key or "" - ) - if ctx: - # Persist ONLY portal-derived values: an OR-fallback value cached - # on a portal blip would be frozen in by step 1 forever. - if base_url and source == "portal": - save_context_length(model, base_url, ctx) - return ctx - if effective_provider == "openai-codex": - # Codex OAuth enforces lower limits than the direct API for the same - # slug (gpt-5.5: 1.05M vs 272K); its own /models is authoritative. - codex_ctx, codex_source = _resolve_codex_oauth_context_length_with_source( - model, access_token=api_key or "", - ) - if codex_ctx: - # Only a fresh authenticated catalogue response may be persisted; - # the static fallback must not poison future probes. - if base_url and codex_source == "live": - save_context_length(model, base_url, codex_ctx) - return codex_ctx - if effective_provider == "gmi" and base_url: - # GMI exposes authoritative context_length via /models, but it is not - # in models.dev yet. Preserve that higher-fidelity endpoint lookup. - ctx = _resolve_endpoint_context_length(model, base_url, api_key=api_key) - if ctx is not None: - return ctx - # 5e. Ollama native /api/show for any base_url that is not a known - # non-Ollama provider (OpenAI-compat /v1/models omits context_length; the - # GGUF model_info is authoritative). Known hosted providers are skipped: - # the POST always 404s and cost ~300ms on the first-turn critical path. - if base_url: - _inferred_for_probe = _infer_provider_from_url(base_url) - _skip_ollama_probe = ( - _inferred_for_probe is not None - and "ollama" not in _inferred_for_probe - ) - if not _skip_ollama_probe: - ctx = _query_ollama_api_show(model, base_url, api_key=api_key) - if ctx is not None: - if not _skip_persistent_context_cache(base_url, provider): - save_context_length(model, base_url, ctx) - return ctx - # 5f. OpenRouter live /models — authoritative for OR-routed models and - # refreshed as new slugs ship, so it must win over models.dev (5g) and the - # family catch-all (8): otherwise a brand-new slug (claude-fable-5, 1M) - # falls through to the generic "claude": 200K entry. - if effective_provider == "openrouter": - metadata = fetch_model_metadata() - entry = metadata.get(model) - if entry: - or_ctx = entry.get("context_length") - # Guard against the known OpenRouter Kimi-family 32k underreport - # (same class the hardcoded overrides exist to mitigate). - if isinstance(or_ctx, int) and or_ctx > 0 and not ( - or_ctx == 32768 and _model_name_suggests_kimi(model) - ): - return or_ctx - - if effective_provider: - from agent.models_dev import lookup_models_dev_context - ctx = lookup_models_dev_context(effective_provider, model) - if ctx: - # MiniMax M3: models.dev reports 512K but actual context is 1M. - # Prefer hardcoded catalog over stale probe value. - if _model_name_suggests_minimax_m3(model): - catalog = DEFAULT_CONTEXT_LENGTHS.get("minimax-m3") - if catalog and ctx < catalog: - logger.info( - "Rejecting models.dev context=%s for %r " - "(MiniMax-M3 underreport); using hardcoded default %s", - ctx, model, f"{catalog:,}", - ) - ctx = catalog - return ctx + ctx = _resolve_provider_aware_context_length(model, base_url, api_key, provider, effective_provider) + if ctx is not None: + return ctx # 6. OpenRouter metadata, provider-unaware fallback — only when the # provider is unknown (OR data is community-maintained). @@ -2858,20 +2678,8 @@ async def get_model_context_length_async( ) -def _is_cjk_token_dense_char(ch: str) -> bool: - code = ord(ch) - return ( - 0x1100 <= code <= 0x11FF # Hangul Jamo - or 0x2E80 <= code <= 0x9FFF # CJK radicals/ideographs - or 0xA960 <= code <= 0xA97F # Hangul Jamo Extended-A - or 0xAC00 <= code <= 0xD7AF # Hangul Syllables - or 0xF900 <= code <= 0xFAFF # CJK compatibility ideographs - or 0xFF00 <= code <= 0xFFEF # Fullwidth forms / halfwidth kana - ) - - -# Same ranges as _is_cjk_token_dense_char (MUST stay in sync) so dense-char -# counting runs in C rather than a per-char Python loop. +# CJK/Hangul/Kana codepoints estimate ~1 token each; a single C-level regex +# pass keeps dense-char counting out of a per-char Python loop. _CJK_DENSE_RE = re.compile( "[\u1100-\u11ff" # Hangul Jamo "\u2e80-\u9fff" # CJK radicals/ideographs @@ -2882,6 +2690,10 @@ _CJK_DENSE_RE = re.compile( ) +def _is_cjk_token_dense_char(ch: str) -> bool: + return _CJK_DENSE_RE.fullmatch(ch) is not None + + def estimate_tokens_rough(text: str) -> int: """Rough token estimate: ceil(chars/4), CJK/Hangul/Kana codepoints ~1 token each. @@ -2898,8 +2710,7 @@ def estimate_tokens_rough(text: str) -> int: dense = len(text) - len(_CJK_DENSE_RE.sub("", text)) if not dense: # non-ASCII but no CJK (accents, Cyrillic, emoji) return (len(text) + 3) // 4 - sparse = len(text) - dense - return dense + ((sparse + 3) // 4) + return dense + ((len(text) - dense + 3) // 4) def estimate_messages_tokens_rough( @@ -2919,16 +2730,12 @@ def estimate_messages_tokens_rough( nothing to compact. Default True is the conservative full charge. Per-message results are memoized on an identity fingerprint (see - ``_estimate_message_tokens_cached``); equal fingerprints imply identical - leaves and structure, hence identical estimates. + ``_estimate_message_tokens_cached``). """ _IMAGE_TOKEN_COST = 1500 if not charge_stale_thinking: messages = _strip_stale_thinking_for_estimate(messages) - total = 0 - for msg in messages: - total += _estimate_message_tokens_cached(msg, _IMAGE_TOKEN_COST) - return total + return sum(_estimate_message_tokens_cached(msg, _IMAGE_TOKEN_COST) for msg in messages) # Generic thinking-text keys replayed for at most the newest assistant turn @@ -2937,20 +2744,17 @@ def estimate_messages_tokens_rough( _STALE_THINKING_ESTIMATE_KEYS = ("reasoning", "reasoning_content") -def _strip_stale_thinking_for_estimate( - messages: List[Dict[str, Any]], -) -> List[Dict[str, Any]]: +def _strip_stale_thinking_for_estimate(messages: List[Dict[str, Any]]) -> List[Dict[str, Any]]: """Copy of ``messages`` with stale thinking keys removed (newest kept). Shallow stripped copies share the original value objects, so the per-message memo still hits for the stripped shape on subsequent walks. """ - newest = -1 - for i in range(len(messages) - 1, -1, -1): - m = messages[i] - if isinstance(m, dict) and m.get("role") == "assistant": - newest = i - break + newest = next( + (i for i in range(len(messages) - 1, -1, -1) + if isinstance(messages[i], dict) and messages[i].get("role") == "assistant"), + -1, + ) out: List[Dict[str, Any]] = [] for i, m in enumerate(messages): if ( @@ -2959,10 +2763,7 @@ def _strip_stale_thinking_for_estimate( and m.get("role") == "assistant" and any(m.get(k) for k in _STALE_THINKING_ESTIMATE_KEYS) ): - m = { - k: v for k, v in m.items() - if k not in _STALE_THINKING_ESTIMATE_KEYS - } + m = {k: v for k, v in m.items() if k not in _STALE_THINKING_ESTIMATE_KEYS} out.append(m) return out @@ -2990,10 +2791,7 @@ def _msg_fingerprint(value: Any, pins: list) -> Any: if t is int or t is float: return ("n", t.__name__, value) if t is dict: - return ("d", tuple( - (_msg_fingerprint(k, pins), _msg_fingerprint(v, pins)) - for k, v in value.items() - )) + return ("d", tuple((_msg_fingerprint(k, pins), _msg_fingerprint(v, pins)) for k, v in value.items())) if t is list: return ("l", tuple(_msg_fingerprint(v, pins) for v in value)) if t is tuple: @@ -3002,22 +2800,19 @@ def _msg_fingerprint(value: Any, pins: list) -> Any: def _estimate_message_tokens_cached(msg: Any, image_cost: int) -> int: + def _compute() -> int: + return _estimate_message_tokens_without_images(msg) + _count_image_tokens(msg, image_cost) + try: pins: list = [] key = _msg_fingerprint(msg, pins) hash(key) except Exception: - return ( - _estimate_message_tokens_without_images(msg) - + _count_image_tokens(msg, image_cost) - ) + return _compute() cached = _MSG_TOKENS_CACHE.get(key) if cached is not None: return cached[1] - tokens = ( - _estimate_message_tokens_without_images(msg) - + _count_image_tokens(msg, image_cost) - ) + tokens = _compute() _MSG_TOKENS_CACHE[key] = (pins, tokens) while len(_MSG_TOKENS_CACHE) > _MSG_TOKENS_CACHE_MAX: try: @@ -3027,29 +2822,22 @@ def _estimate_message_tokens_cached(msg: Any, image_cost: int) -> int: return tokens +def _count_parts(parts: Any, types: set) -> int: + if not isinstance(parts, list): + return 0 + return sum(1 for part in parts if isinstance(part, dict) and part.get("type") in types) + + def _count_image_tokens(msg: Dict[str, Any], cost_per_image: int) -> int: """Count image-like content parts in a message; return their token cost.""" - count = 0 - content = msg.get("content") if isinstance(msg, dict) else None - if isinstance(content, list): - for part in content: - if not isinstance(part, dict): - continue - ptype = part.get("type") - if ptype in {"image", "image_url", "input_image"}: - count += 1 - stashed = msg.get("_anthropic_content_blocks") if isinstance(msg, dict) else None - if isinstance(stashed, list): - for part in stashed: - if isinstance(part, dict) and part.get("type") == "image": - count += 1 + if not isinstance(msg, dict): + return 0 + content = msg.get("content") + count = _count_parts(content, {"image", "image_url", "input_image"}) + count += _count_parts(msg.get("_anthropic_content_blocks"), {"image"}) # Multimodal tool results that haven't been converted yet. if isinstance(content, dict) and content.get("_multimodal"): - inner = content.get("content") - if isinstance(inner, list): - for part in inner: - if isinstance(part, dict) and part.get("type") in {"image", "image_url"}: - count += 1 + count += _count_parts(content.get("content"), {"image", "image_url"}) return count * cost_per_image @@ -3063,18 +2851,13 @@ def _wire_message_shadow(msg: Dict[str, Any]) -> Dict[str, Any]: the wire, and substituting it would UNDERcount — the dangerous direction (compaction fires too late, the turn dies on a hard context error). * Base64 images become a placeholder; ``_count_image_tokens`` charges them flat. + * ``reasoning`` never ships as-is: request builds pop it after optionally + promoting it into ``reasoning_content``. When both exist counting both + inflated estimates up to +53%; keep ``reasoning`` only as the promotion + proxy when nothing displaces it. """ sidecar = msg.get("api_content") - sidecar_wins = ( - isinstance(sidecar, str) - and bool(sidecar) - and msg.get("role") in ("user", "assistant") - ) - # ``reasoning`` never ships as-is: request builds pop it after optionally - # promoting it into ``reasoning_content``. When both exist (reasoning-echo - # providers pin reasoning_content; reasoning holds the same text for the - # trajectory) counting both inflated estimates up to +53%; keep - # ``reasoning`` only as the promotion proxy when nothing displaces it. + sidecar_wins = isinstance(sidecar, str) and bool(sidecar) and msg.get("role") in ("user", "assistant") _rc = msg.get("reasoning_content") drop_reasoning_dup = isinstance(_rc, str) and bool(_rc.strip()) shadow: Dict[str, Any] = {} @@ -3086,21 +2869,15 @@ def _wire_message_shadow(msg: Dict[str, Any]) -> Dict[str, Any]: if k == "api_content": if sidecar_wins: shadow["content"] = v - continue - if k == "content": + elif k == "content": if sidecar_wins: continue if isinstance(v, list): - cleaned = [] - for part in v: - if isinstance(part, dict): - if part.get("type") in {"image", "image_url", "input_image"}: - cleaned.append({"type": part.get("type"), "image": "[stripped]"}) - else: - cleaned.append(part) - else: - cleaned.append(part) - shadow[k] = cleaned + shadow[k] = [ + {"type": part.get("type"), "image": "[stripped]"} + if isinstance(part, dict) and part.get("type") in {"image", "image_url", "input_image"} else part + for part in v + ] elif isinstance(v, dict) and v.get("_multimodal"): shadow[k] = v.get("text_summary", "") else: @@ -3137,9 +2914,7 @@ def estimate_request_tokens_rough( # estimate_messages_tokens_rough with (messages)-only signatures. total += estimate_messages_tokens_rough(messages) else: - total += estimate_messages_tokens_rough( - messages, charge_stale_thinking=False - ) + total += estimate_messages_tokens_rough(messages, charge_stale_thinking=False) if tools: total += _estimate_tools_tokens_rough(tools) return total @@ -3170,12 +2945,11 @@ def capture_usage_anchor( return None if pt <= 0 or not isinstance(messages, list): return None # no usable usage (some endpoints omit it) — caller keeps its anchor - base_count = len(messages) - last = messages[-1] if base_count else None + last = messages[-1] if messages else None return { "prompt_tokens": pt, "completion_tokens": max(0, ct), - "base_count": base_count, + "base_count": len(messages), "base_last_id": id(last) if last is not None else None, "base_last_role": last.get("role") if isinstance(last, dict) else None, } @@ -3204,14 +2978,10 @@ def anchored_context_tokens( return None total = int(anchor["prompt_tokens"]) + int(anchor.get("completion_tokens") or 0) delta = messages[base_count:] + if delta and isinstance(delta[0], dict) and delta[0].get("role") == "assistant": + delta = delta[1:] if delta: - first = delta[0] - if isinstance(first, dict) and first.get("role") == "assistant": - delta = delta[1:] - if delta: - total += estimate_messages_tokens_rough( - delta, charge_stale_thinking=charge_stale_thinking - ) + total += estimate_messages_tokens_rough(delta, charge_stale_thinking=charge_stale_thinking) return total @@ -3225,10 +2995,9 @@ def _tool_name_for_cache(tool: Any) -> str: if not isinstance(tool, dict): return "" fn = tool.get("function") - if isinstance(fn, dict): - name = fn.get("name") - if isinstance(name, str): - return name + name = fn.get("name") if isinstance(fn, dict) else None + if isinstance(name, str): + return name name = tool.get("name") return name if isinstance(name, str) else "" @@ -3239,14 +3008,12 @@ def _estimate_tools_tokens_rough(tools: List[Dict[str, Any]]) -> int: key = id(tools) n = len(tools) - first = _tool_name_for_cache(tools[0]) if n else "" - last = _tool_name_for_cache(tools[-1]) if n else "" + first = _tool_name_for_cache(tools[0]) + last = _tool_name_for_cache(tools[-1]) cached = _TOOLS_TOKENS_CACHE.get(key) - if cached is not None: - cached_n, cached_first, cached_last, cached_tokens = cached - if cached_n == n and cached_first == first and cached_last == last: - return cached_tokens + if cached is not None and cached[:3] == (n, first, last): + return cached[3] # Sum the major schema fields (descriptions + parameters dominate). total_chars = 0 @@ -3254,15 +3021,10 @@ def _estimate_tools_tokens_rough(tools: List[Dict[str, Any]]) -> int: if not isinstance(tool, dict): continue fn = tool.get("function") - if isinstance(fn, dict): - name = fn.get("name") or "" - desc = fn.get("description") or "" - params = fn.get("parameters") or {} - else: - name = tool.get("name") or "" - desc = tool.get("description") or "" - params = tool.get("parameters") or {} - + src = fn if isinstance(fn, dict) else tool + name = src.get("name") or "" + desc = src.get("description") or "" + params = src.get("parameters") or {} if isinstance(name, str): total_chars += len(name) if isinstance(desc, str):