refactor(agent/models): dedupe model_metadata probe helpers
- _server_root / _ollama_show_context / _longest_key_match / _probe_local_context_length / _endpoint_model_entry replace 5 copies of the same local-probe and lookup bodies - detect_local_server_type waterfall as an ordered (name, probe) table - output-cap error classification via phrase-group tables shared by is_output_cap_error and parse_available_output_tokens_from_error - endpoint-scoped context overrides as a data table - drop dead _fetch_codex_oauth_context_lengths, _resolve_codex_oauth_context_length, _estimate_message_chars (zero refs repo-wide) Verified with a differential harness (old vs new module, ~700 pure-function probes, 0 mismatches).
This commit is contained in:
+243
-386
@@ -500,7 +500,7 @@ DEFAULT_CONTEXT_LENGTHS = {
|
||||
# Source: https://developers.openai.com/api/docs/models
|
||||
# GPT-5.5 (launched Apr 23 2026) is 1.05M on the direct OpenAI API and
|
||||
# ChatGPT Codex OAuth caps it at 272K; both paths resolve via their own
|
||||
# provider-aware branches (_resolve_codex_oauth_context_length + models.dev).
|
||||
# provider-aware branches (_resolve_codex_oauth_context_length_with_source + models.dev).
|
||||
# This hardcoded value is only reached when every probe misses.
|
||||
# GPT-5.6 series (Sol/Terra/Luna, GA 2026-07-09) — 1.05M on the direct
|
||||
# OpenAI API (same as gpt-5.5). Codex OAuth caps these at 272K.
|
||||
@@ -864,6 +864,70 @@ def _is_known_provider_base_url(base_url: str) -> bool:
|
||||
return _infer_provider_from_url(base_url) is not None
|
||||
|
||||
|
||||
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
|
||||
|
||||
|
||||
def _longest_key_match(table: Dict[str, int], model_lower: str) -> Optional[Tuple[str, int]]:
|
||||
"""First ``(key, value)`` whose key is a substring of ``model_lower``, longest key first.
|
||||
|
||||
Longest-first makes specific entries (``gpt-5.4-mini``) win over their
|
||||
family catch-all (``gpt-5``); ties keep table order (stable sort).
|
||||
"""
|
||||
for key, value in sorted(table.items(), key=lambda x: len(x[0]), reverse=True):
|
||||
if key in model_lower:
|
||||
return key, value
|
||||
return None
|
||||
|
||||
|
||||
def _ollama_show_context(data: Dict[str, Any], *, gguf_first: bool, minimum: Optional[int] = None) -> Optional[int]:
|
||||
"""Context length from an Ollama ``/api/show`` payload.
|
||||
|
||||
``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``).
|
||||
"""
|
||||
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
|
||||
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
|
||||
return None
|
||||
|
||||
for reader in ((_gguf, _num_ctx) if gguf_first else (_num_ctx, _gguf)):
|
||||
ctx = reader()
|
||||
if ctx is not None:
|
||||
return ctx
|
||||
return None
|
||||
|
||||
|
||||
# (host, canonical paths, model ids, context) — see _endpoint_scoped_context_length.
|
||||
_ENDPOINT_SCOPED_CONTEXT = (
|
||||
("api.kimi.com", {"/coding", "/coding/v1"}, {"k3", "kimi-k3", "kimi-k3-cot"}, 1_048_576),
|
||||
("integrate.api.nvidia.com", {"/v1"}, {"deepseek-ai/deepseek-v4-pro"}, 262_144),
|
||||
)
|
||||
|
||||
|
||||
def _endpoint_scoped_context_length(model: str, base_url: str) -> Optional[int]:
|
||||
"""Return context metadata confirmed for one provider endpoint.
|
||||
|
||||
@@ -883,30 +947,22 @@ def _endpoint_scoped_context_length(model: str, base_url: str) -> Optional[int]:
|
||||
port = parsed.port
|
||||
except ValueError:
|
||||
return None
|
||||
# Only canonical https://host[:443]/path with no credentials/query/fragment.
|
||||
if (
|
||||
parsed.scheme.lower() == "https"
|
||||
and (parsed.hostname or "").lower() == "api.kimi.com"
|
||||
and port in (None, 443)
|
||||
and parsed.username is None
|
||||
and parsed.password is None
|
||||
and parsed.path.rstrip("/") in {"/coding", "/coding/v1"}
|
||||
and not parsed.query
|
||||
and not parsed.fragment
|
||||
and model.strip().lower() in {"k3", "kimi-k3", "kimi-k3-cot"}
|
||||
parsed.scheme.lower() != "https"
|
||||
or port not in (None, 443)
|
||||
or parsed.username is not None
|
||||
or parsed.password is not None
|
||||
or parsed.query
|
||||
or parsed.fragment
|
||||
):
|
||||
return 1_048_576
|
||||
if (
|
||||
parsed.scheme.lower() == "https"
|
||||
and (parsed.hostname or "").lower() == "integrate.api.nvidia.com"
|
||||
and port in (None, 443)
|
||||
and parsed.username is None
|
||||
and parsed.password is None
|
||||
and parsed.path.rstrip("/") == "/v1"
|
||||
and not parsed.query
|
||||
and not parsed.fragment
|
||||
and model.strip().lower() == "deepseek-ai/deepseek-v4-pro"
|
||||
):
|
||||
return 262_144
|
||||
return None
|
||||
host = (parsed.hostname or "").lower()
|
||||
path = parsed.path.rstrip("/")
|
||||
model_key = model.strip().lower()
|
||||
for scoped_host, paths, models, ctx in _ENDPOINT_SCOPED_CONTEXT:
|
||||
if host == scoped_host and path in paths and model_key in models:
|
||||
return ctx
|
||||
return None
|
||||
|
||||
|
||||
@@ -940,6 +996,16 @@ def _maybe_cache_local_context_length(
|
||||
save_context_length(model, base_url, length)
|
||||
|
||||
|
||||
def _probe_local_context_length(model: str, base_url: str, api_key: str, provider: str) -> Optional[int]:
|
||||
"""Live local probe; persists a positive result unless the provider opts out of the disk cache."""
|
||||
local_ctx = _query_local_context_length(model, base_url, api_key=api_key)
|
||||
if local_ctx and local_ctx > 0:
|
||||
if not _skip_persistent_context_cache(base_url, provider):
|
||||
_maybe_cache_local_context_length(model, base_url, local_ctx)
|
||||
return local_ctx
|
||||
return None
|
||||
|
||||
|
||||
def _reconcile_local_cached_context_length(
|
||||
model: str,
|
||||
base_url: str,
|
||||
@@ -1020,16 +1086,14 @@ def is_local_endpoint(base_url: str) -> bool:
|
||||
if len(parts) == 4:
|
||||
try:
|
||||
first, second = int(parts[0]), int(parts[1])
|
||||
if first == 10:
|
||||
return True
|
||||
if first == 172 and 16 <= second <= 31:
|
||||
return True
|
||||
if first == 192 and second == 168:
|
||||
return True
|
||||
if first == 100 and 64 <= second <= 127:
|
||||
return True
|
||||
except ValueError:
|
||||
pass
|
||||
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
|
||||
|
||||
|
||||
@@ -1070,16 +1134,10 @@ def detect_local_server_type(base_url: str, api_key: str = "") -> Optional[str]:
|
||||
"""
|
||||
import httpx
|
||||
|
||||
normalized = _normalize_base_url(base_url)
|
||||
|
||||
# Resolve localhost to IPv4 to avoid 2s IPv6 timeout on Windows dual-stack.
|
||||
# Applied to ``normalized`` before deriving server/LM Studio URLs AND
|
||||
# before the cache lookup, so localhost and 127.0.0.1 share a cache entry.
|
||||
normalized = _localhost_to_ipv4(normalized)
|
||||
|
||||
server_url = normalized
|
||||
if server_url.endswith("/v1"):
|
||||
server_url = server_url[:-3]
|
||||
# IPv4-resolve BEFORE deriving server/LM Studio URLs and the cache lookup,
|
||||
# so localhost and 127.0.0.1 share a cache entry.
|
||||
normalized = _localhost_to_ipv4(_normalize_base_url(base_url))
|
||||
server_url = _server_root(normalized)
|
||||
lmstudio_url = _lmstudio_server_root(normalized)
|
||||
|
||||
cached = _endpoint_probe_path_cache.get(server_url)
|
||||
@@ -1122,49 +1180,35 @@ def detect_local_server_type(base_url: str, api_key: str = "") -> Optional[str]:
|
||||
_note_endpoint_blackholed(server_url)
|
||||
raise exc
|
||||
|
||||
def _lm_studio(client) -> bool:
|
||||
return client.get(f"{lmstudio_url}/api/v1/models").status_code == 200
|
||||
|
||||
def _ollama(client) -> bool:
|
||||
# LM Studio answers /api/tags with {"error": ...} and status 200, so
|
||||
# the body must actually carry "models".
|
||||
r = client.get(f"{server_url}/api/tags")
|
||||
return r.status_code == 200 and "models" in r.json()
|
||||
|
||||
def _llamacpp(client) -> bool:
|
||||
r = client.get(f"{server_url}/v1/props")
|
||||
if r.status_code != 200:
|
||||
r = client.get(f"{server_url}/props") # older builds: no /v1 prefix
|
||||
return r.status_code == 200 and "default_generation_settings" in r.text
|
||||
|
||||
def _vllm(client) -> bool:
|
||||
r = client.get(f"{server_url}/version")
|
||||
return r.status_code == 200 and "version" in r.json()
|
||||
|
||||
# Most specific first: LM Studio's native API, then Ollama, llama.cpp, vLLM.
|
||||
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:
|
||||
# LM Studio exposes /api/v1/models — check first (most specific)
|
||||
try:
|
||||
r = client.get(f"{lmstudio_url}/api/v1/models")
|
||||
if r.status_code == 200:
|
||||
result = "lm-studio"
|
||||
except Exception as exc:
|
||||
_probe_failed(exc)
|
||||
if result is None:
|
||||
# Ollama exposes /api/tags and responds with {"models": [...]}
|
||||
# LM Studio returns {"error": "Unexpected endpoint"} with status 200
|
||||
# on this path, so we must verify the response contains "models".
|
||||
for name, probe in waterfall:
|
||||
try:
|
||||
r = client.get(f"{server_url}/api/tags")
|
||||
if r.status_code == 200:
|
||||
try:
|
||||
data = r.json()
|
||||
if "models" in data:
|
||||
result = "ollama"
|
||||
except Exception:
|
||||
pass
|
||||
except Exception as exc:
|
||||
_probe_failed(exc)
|
||||
if result is None:
|
||||
# llama.cpp exposes /v1/props (older builds used /props without the /v1 prefix)
|
||||
try:
|
||||
r = client.get(f"{server_url}/v1/props")
|
||||
if r.status_code != 200:
|
||||
r = client.get(f"{server_url}/props") # fallback for older builds
|
||||
if r.status_code == 200 and "default_generation_settings" in r.text:
|
||||
result = "llamacpp"
|
||||
except Exception as exc:
|
||||
_probe_failed(exc)
|
||||
if result is None:
|
||||
# vLLM: /version
|
||||
try:
|
||||
r = client.get(f"{server_url}/version")
|
||||
if r.status_code == 200:
|
||||
data = r.json()
|
||||
if "version" in data:
|
||||
result = "vllm"
|
||||
if probe(client):
|
||||
result = name
|
||||
break
|
||||
except Exception as exc:
|
||||
_probe_failed(exc)
|
||||
except Exception:
|
||||
@@ -1276,33 +1320,25 @@ def _context_length_from_model_payload(payload: Dict[str, Any]) -> Optional[int]
|
||||
|
||||
|
||||
def _extract_pricing(payload: Dict[str, Any]) -> Dict[str, Any]:
|
||||
novita_input = payload.get("input_token_price_per_m")
|
||||
novita_output = payload.get("output_token_price_per_m")
|
||||
if novita_input is not None or novita_output is not None:
|
||||
pricing: Dict[str, Any] = {}
|
||||
if novita_input is not None:
|
||||
pricing["prompt"] = str(float(novita_input) / 10_000 / 1_000_000)
|
||||
if novita_output is not None:
|
||||
pricing["completion"] = str(float(novita_output) / 10_000 / 1_000_000)
|
||||
return pricing
|
||||
def _per_token(source: Dict[str, Any], fields: Dict[str, str], scale) -> Dict[str, Any]:
|
||||
# Provider $/MTok (or Novita's 1/10_000-$ per M) -> per-token strings so
|
||||
# usage_pricing consumes them through the same path as OpenRouter.
|
||||
return {
|
||||
target: str(scale(float(source[key])))
|
||||
for target, key in fields.items()
|
||||
if source.get(key) is not None
|
||||
}
|
||||
|
||||
# DeepInfra ships pricing under ``metadata.pricing`` with $/MTok values:
|
||||
# ``input_tokens``, ``output_tokens``, ``cache_read_tokens``. Convert to
|
||||
# per-token strings so the generic cost machinery (usage_pricing.py)
|
||||
# consumes them through the same path as OpenRouter / OpenAI.
|
||||
novita_fields = {"prompt": "input_token_price_per_m", "completion": "output_token_price_per_m"}
|
||||
if any(payload.get(k) is not None for k in novita_fields.values()):
|
||||
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
|
||||
if isinstance(deepinfra_pricing, dict) and any(
|
||||
k in deepinfra_pricing for k in ("input_tokens", "output_tokens", "cache_read_tokens")
|
||||
):
|
||||
result: Dict[str, Any] = {}
|
||||
if deepinfra_pricing.get("input_tokens") is not None:
|
||||
result["prompt"] = str(float(deepinfra_pricing["input_tokens"]) / 1_000_000)
|
||||
if deepinfra_pricing.get("output_tokens") is not None:
|
||||
result["completion"] = str(float(deepinfra_pricing["output_tokens"]) / 1_000_000)
|
||||
if deepinfra_pricing.get("cache_read_tokens") is not None:
|
||||
result["cache_read"] = str(float(deepinfra_pricing["cache_read_tokens"]) / 1_000_000)
|
||||
return result
|
||||
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)
|
||||
|
||||
alias_map = {
|
||||
"prompt": ("prompt", "input", "input_cost_per_token", "prompt_token_cost"),
|
||||
@@ -1394,6 +1430,20 @@ def fetch_model_metadata(force_refresh: bool = False) -> Dict[str, Dict[str, Any
|
||||
return {}
|
||||
|
||||
|
||||
def _endpoint_model_entry(model: Dict[str, Any], model_id: str, context_length: Optional[int]) -> Dict[str, Any]:
|
||||
"""Cache entry for one ``/models`` item; optional keys are set only when known."""
|
||||
entry: Dict[str, Any] = {"name": model.get("name", model_id)}
|
||||
if context_length is not None:
|
||||
entry["context_length"] = context_length
|
||||
max_completion_tokens = _extract_max_completion_tokens(model)
|
||||
if max_completion_tokens is not None:
|
||||
entry["max_completion_tokens"] = max_completion_tokens
|
||||
pricing = _extract_pricing(model)
|
||||
if pricing:
|
||||
entry["pricing"] = pricing
|
||||
return entry
|
||||
|
||||
|
||||
def fetch_endpoint_model_metadata(
|
||||
base_url: str,
|
||||
api_key: str = "",
|
||||
@@ -1457,8 +1507,6 @@ def fetch_endpoint_model_metadata(
|
||||
model_id = model.get("key") or model.get("id")
|
||||
if not model_id:
|
||||
continue
|
||||
entry: Dict[str, Any] = {"name": model.get("name", model_id)}
|
||||
|
||||
context_length = None
|
||||
for inst in model.get("loaded_instances", []) or []:
|
||||
if not isinstance(inst, dict):
|
||||
@@ -1468,17 +1516,7 @@ def fetch_endpoint_model_metadata(
|
||||
if isinstance(ctx, int) and ctx > 0:
|
||||
context_length = ctx
|
||||
break
|
||||
if context_length is not None:
|
||||
entry["context_length"] = context_length
|
||||
|
||||
max_completion_tokens = _extract_max_completion_tokens(model)
|
||||
if max_completion_tokens is not None:
|
||||
entry["max_completion_tokens"] = max_completion_tokens
|
||||
|
||||
pricing = _extract_pricing(model)
|
||||
if pricing:
|
||||
entry["pricing"] = pricing
|
||||
|
||||
entry = _endpoint_model_entry(model, model_id, context_length)
|
||||
_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:
|
||||
@@ -1528,17 +1566,7 @@ def fetch_endpoint_model_metadata(
|
||||
model_id = model.get("id")
|
||||
if not model_id:
|
||||
continue
|
||||
entry: Dict[str, Any] = {"name": model.get("name", model_id)}
|
||||
context_length = _extract_context_length(model)
|
||||
if context_length is not None:
|
||||
entry["context_length"] = context_length
|
||||
max_completion_tokens = _extract_max_completion_tokens(model)
|
||||
if max_completion_tokens is not None:
|
||||
entry["max_completion_tokens"] = max_completion_tokens
|
||||
pricing = _extract_pricing(model)
|
||||
if pricing:
|
||||
entry["pricing"] = pricing
|
||||
_add_model_aliases(cache, model_id, entry)
|
||||
_add_model_aliases(cache, model_id, _endpoint_model_entry(model, model_id, _extract_context_length(model)))
|
||||
|
||||
# If this is a llama.cpp server, query /props for actual allocated context
|
||||
is_llamacpp = any(
|
||||
@@ -1835,80 +1863,24 @@ def parse_available_output_tokens_from_error(error_msg: str) -> Optional[int]:
|
||||
the error does not look like a max_tokens-too-large error.
|
||||
"""
|
||||
error_lower = error_msg.lower()
|
||||
|
||||
# Must look like an output-cap error, not a prompt-length error.
|
||||
is_output_cap_error = (
|
||||
"max_tokens" in error_lower
|
||||
and ("available_tokens" in error_lower or "available tokens" in error_lower)
|
||||
) or (
|
||||
# OpenRouter/Nous phrasing of the same condition.
|
||||
"in the output" in error_lower
|
||||
and "maximum context length" in error_lower
|
||||
) or (
|
||||
# LM Studio / llama.cpp / some OpenAI-compatible servers:
|
||||
# "This model's maximum context length is 65536 tokens. However, you
|
||||
# requested 65536 output tokens and your prompt contains 77409
|
||||
# characters ..."
|
||||
# The "requested N output tokens" phrasing means the OUTPUT cap is the
|
||||
# problem (the input itself fits) — reduce max_tokens, don't compress.
|
||||
"maximum context length" in error_lower
|
||||
and "requested" in error_lower
|
||||
and "output tokens" in error_lower
|
||||
) or (
|
||||
# DashScope / Alibaba Cloud (Qwen) phrasing. The provider rejects an
|
||||
# over-cap output request with a bounded range whose upper bound IS the
|
||||
# real max-output cap, e.g.
|
||||
# "Range of max_tokens should be [1, 65536]"
|
||||
# The input itself fits — this is purely an output-cap error, so reduce
|
||||
# max_tokens and retry; do NOT compress.
|
||||
"range of max_tokens should be" in error_lower
|
||||
) or (
|
||||
# OpenAI-compatible relays may reject a request whose output cap exceeds
|
||||
# the model's separate completion-token limit, e.g.
|
||||
# "max_tokens (98304) exceeds model's maximum output tokens (65536)"
|
||||
# This is independent of the input context window.
|
||||
"exceeds model" in error_lower
|
||||
and "maximum output tokens" in error_lower
|
||||
)
|
||||
if not is_output_cap_error:
|
||||
if not _any_phrase_group(error_lower, _PARSEABLE_OUTPUT_CAP_SIGNALS):
|
||||
return None
|
||||
|
||||
# Generic model-output-cap form:
|
||||
# Direct cap figures, most specific first:
|
||||
# "max_tokens (98304) exceeds model's maximum output tokens (65536)"
|
||||
_m_max_output = re.search(
|
||||
# "Range of max_tokens should be [1, 65536]" (upper bound is the cap)
|
||||
# "... = available_tokens: 10000" (Anthropic)
|
||||
# "200000 - 190000 = 10000" (last number after "=")
|
||||
for pattern in (
|
||||
r'exceeds model(?:\'s)? maximum output tokens\s*\(?\s*(\d+)\s*\)?',
|
||||
error_lower,
|
||||
)
|
||||
if _m_max_output:
|
||||
_cap = int(_m_max_output.group(1))
|
||||
if _cap >= 1:
|
||||
return _cap
|
||||
|
||||
# DashScope / Alibaba range form: "Range of max_tokens should be [1, 65536]".
|
||||
# The upper bound is the available output cap.
|
||||
_m_range = re.search(
|
||||
r'range of max_tokens should be\s*\[\s*\d+\s*,\s*(\d+)\s*\]',
|
||||
error_lower,
|
||||
)
|
||||
if _m_range:
|
||||
_cap = int(_m_range.group(1))
|
||||
if _cap >= 1:
|
||||
return _cap
|
||||
|
||||
# Extract the available_tokens figure.
|
||||
# Anthropic format: "… = available_tokens: 10000"
|
||||
patterns = [
|
||||
r'available_tokens[:\s]+(\d+)',
|
||||
r'available\s+tokens[:\s]+(\d+)',
|
||||
# fallback: last number after "=" in expressions like "200000 - 190000 = 10000"
|
||||
r'=\s*(\d+)\s*$',
|
||||
]
|
||||
for pattern in patterns:
|
||||
):
|
||||
match = re.search(pattern, error_lower)
|
||||
if match:
|
||||
tokens = int(match.group(1))
|
||||
if tokens >= 1:
|
||||
return tokens
|
||||
if match and int(match.group(1)) >= 1:
|
||||
return int(match.group(1))
|
||||
|
||||
# OpenRouter/Nous format: "maximum context length is N … (A of text input,
|
||||
# B of tool input, C in the output)". Available output = ctx - text - tool.
|
||||
@@ -1975,6 +1947,41 @@ def parse_available_output_tokens_from_error(error_msg: str) -> Optional[int]:
|
||||
return None
|
||||
|
||||
|
||||
# Each entry is a phrase group; the group matches when ALL phrases are present.
|
||||
_OUTPUT_CAP_SIGNALS = (
|
||||
("range of max_tokens should be",), # DashScope / Alibaba
|
||||
("available_tokens",), # Anthropic
|
||||
("available tokens",),
|
||||
("in the output", "maximum context length"), # OpenRouter / Nous
|
||||
("requested", "output tokens"), # LM Studio / llama.cpp
|
||||
("should be",), # generic "max_tokens should be <= N"
|
||||
("less than or equal",),
|
||||
("must be",),
|
||||
("exceeds model", "maximum output tokens"), # OpenAI-compatible relays
|
||||
)
|
||||
_INPUT_OVERFLOW_SIGNALS = (
|
||||
"prompt is too long", "prompt too long", "input is too long", "input token",
|
||||
"prompt length", "prompt contains", "reduce the length",
|
||||
)
|
||||
# Narrower than _OUTPUT_CAP_SIGNALS: only phrasings we can extract a number from.
|
||||
_PARSEABLE_OUTPUT_CAP_SIGNALS = (
|
||||
("max_tokens", "available_tokens"), # Anthropic
|
||||
("max_tokens", "available tokens"),
|
||||
("in the output", "maximum context length"), # OpenRouter / Nous
|
||||
# "requested N output tokens" means the OUTPUT cap is the problem (the
|
||||
# input itself fits) — reduce max_tokens, don't compress.
|
||||
("maximum context length", "requested", "output tokens"), # LM Studio / llama.cpp
|
||||
# DashScope rejects an over-cap output request with a bounded range whose
|
||||
# upper bound IS the real max-output cap: "Range of max_tokens should be [1, 65536]".
|
||||
("range of max_tokens should be",),
|
||||
("exceeds model", "maximum output tokens"), # "max_tokens (98304) exceeds model's maximum output tokens (65536)"
|
||||
)
|
||||
|
||||
|
||||
def _any_phrase_group(text: str, groups: tuple) -> bool:
|
||||
return any(all(p in text for p in group) for group in groups)
|
||||
|
||||
|
||||
def is_output_cap_error(error_msg: str) -> bool:
|
||||
"""Return True if a 400 is about the OUTPUT cap (max_tokens) being too large.
|
||||
|
||||
@@ -1999,46 +2006,13 @@ def is_output_cap_error(error_msg: str) -> bool:
|
||||
path (a real input overflow can also mention max_tokens).
|
||||
"""
|
||||
error_lower = error_msg.lower()
|
||||
|
||||
mentions_output_param = (
|
||||
"max_tokens" in error_lower
|
||||
or "max_output_tokens" in error_lower
|
||||
or "max_completion_tokens" in error_lower
|
||||
)
|
||||
if not mentions_output_param:
|
||||
if not any(p in error_lower for p in ("max_tokens", "max_output_tokens", "max_completion_tokens")):
|
||||
return False
|
||||
|
||||
# Phrasing that signals the OUTPUT cap specifically is the problem.
|
||||
output_cap_signal = (
|
||||
"range of max_tokens should be" in error_lower # DashScope / Alibaba
|
||||
or "available_tokens" in error_lower # Anthropic
|
||||
or "available tokens" in error_lower
|
||||
or ("in the output" in error_lower # OpenRouter / Nous
|
||||
and "maximum context length" in error_lower)
|
||||
or ("requested" in error_lower # LM Studio / llama.cpp
|
||||
and "output tokens" in error_lower)
|
||||
or "should be" in error_lower # generic "max_tokens should be <= N"
|
||||
or "less than or equal" in error_lower
|
||||
or "must be" in error_lower
|
||||
or ("exceeds model" in error_lower
|
||||
and "maximum output tokens" in error_lower)
|
||||
)
|
||||
if not output_cap_signal:
|
||||
if not _any_phrase_group(error_lower, _OUTPUT_CAP_SIGNALS):
|
||||
return False
|
||||
|
||||
# If the error ALSO clearly describes an oversized INPUT, it is a genuine
|
||||
# context overflow that happens to mention max_tokens — let the
|
||||
# context-overflow path handle it (it can compress the input).
|
||||
input_overflow_signal = (
|
||||
"prompt is too long" in error_lower
|
||||
or "prompt too long" in error_lower
|
||||
or "input is too long" in error_lower
|
||||
or "input token" in error_lower
|
||||
or "prompt length" in error_lower
|
||||
or "prompt contains" in error_lower
|
||||
or "reduce the length" in error_lower
|
||||
)
|
||||
return not input_overflow_signal
|
||||
# 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)
|
||||
|
||||
|
||||
def _model_id_matches(candidate_id: str, lookup_model: str) -> bool:
|
||||
@@ -2073,9 +2047,7 @@ def query_ollama_num_ctx(model: str, base_url: str, api_key: str = "") -> Option
|
||||
import httpx
|
||||
|
||||
bare_model = _strip_provider_prefix(model)
|
||||
server_url = _localhost_to_ipv4(base_url.rstrip("/"))
|
||||
if server_url.endswith("/v1"):
|
||||
server_url = server_url[:-3]
|
||||
server_url = _server_root(base_url)
|
||||
|
||||
try:
|
||||
server_type = detect_local_server_type(base_url, api_key=api_key)
|
||||
@@ -2098,29 +2070,10 @@ def query_ollama_num_ctx(model: str, base_url: str, api_key: str = "") -> Option
|
||||
resp = client.post(f"{server_url}/api/show", json={"name": bare_model})
|
||||
if resp.status_code != 200:
|
||||
return None
|
||||
data = resp.json()
|
||||
|
||||
# Prefer explicit num_ctx from Modelfile parameters (user override)
|
||||
params = data.get("parameters", "")
|
||||
if "num_ctx" in params:
|
||||
for line in params.split("\n"):
|
||||
if "num_ctx" in line:
|
||||
parts = line.strip().split()
|
||||
if len(parts) >= 2:
|
||||
try:
|
||||
_ctx = int(parts[-1])
|
||||
_local_probe_disk_put("ollama_num_ctx", _disk_key, _ctx)
|
||||
return _ctx
|
||||
except ValueError:
|
||||
pass
|
||||
|
||||
# Fall back to GGUF model_info context_length (training max)
|
||||
model_info = data.get("model_info", {})
|
||||
for key, value in model_info.items():
|
||||
if "context_length" in key and isinstance(value, (int, float)):
|
||||
_ctx = int(value)
|
||||
_local_probe_disk_put("ollama_num_ctx", _disk_key, _ctx)
|
||||
return _ctx
|
||||
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
|
||||
@@ -2145,10 +2098,7 @@ def query_ollama_supports_vision(model: str, base_url: str, api_key: str = "") -
|
||||
except Exception:
|
||||
return None
|
||||
|
||||
server_url = _localhost_to_ipv4(base_url.rstrip("/"))
|
||||
if server_url.endswith("/v1"):
|
||||
server_url = server_url[:-3]
|
||||
|
||||
server_url = _server_root(base_url)
|
||||
headers = _auth_headers(api_key)
|
||||
|
||||
try:
|
||||
@@ -2222,10 +2172,7 @@ def _query_ollama_api_show_uncached(model: str, base_url: str, api_key: str = ""
|
||||
"""Uncached body of ``_query_ollama_api_show`` — one POST to ``/api/show``."""
|
||||
import httpx
|
||||
|
||||
server_url = _localhost_to_ipv4(base_url.rstrip("/"))
|
||||
if server_url.endswith("/v1"):
|
||||
server_url = server_url[:-3]
|
||||
|
||||
server_url = _server_root(base_url)
|
||||
if _endpoint_blackholed(server_url):
|
||||
return None
|
||||
|
||||
@@ -2236,30 +2183,11 @@ def _query_ollama_api_show_uncached(model: str, base_url: str, api_key: str = ""
|
||||
resp = client.post(f"{server_url}/api/show", json={"name": model})
|
||||
if resp.status_code != 200:
|
||||
return None
|
||||
data = resp.json()
|
||||
|
||||
# Hosted Ollama: GGUF model_info is the real max — prefer it over
|
||||
# num_ctx which the Cloud operator may have capped arbitrarily.
|
||||
model_info = data.get("model_info", {})
|
||||
for key, value in model_info.items():
|
||||
if "context_length" in key and isinstance(value, (int, float)):
|
||||
ctx = int(value)
|
||||
if ctx >= 1024:
|
||||
return ctx
|
||||
|
||||
# Fall back to num_ctx from Modelfile parameters (rare on Cloud)
|
||||
params = data.get("parameters", "")
|
||||
if "num_ctx" in params:
|
||||
for line in params.split("\n"):
|
||||
if "num_ctx" in line:
|
||||
parts = line.strip().split()
|
||||
if len(parts) >= 2:
|
||||
try:
|
||||
ctx = int(parts[-1])
|
||||
if ctx >= 1024:
|
||||
return ctx
|
||||
except ValueError:
|
||||
pass
|
||||
# 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
|
||||
except Exception as exc:
|
||||
if _is_connect_timeout(exc):
|
||||
_note_endpoint_blackholed(server_url)
|
||||
@@ -2400,10 +2328,7 @@ def _query_local_context_length_uncached(model: str, base_url: str, api_key: str
|
||||
# Ollama "model:tag" colons (e.g. "qwen3.5:27b") are intentionally preserved.
|
||||
model = _strip_provider_prefix(model)
|
||||
|
||||
# Strip /v1 suffix to get the server root
|
||||
server_url = _localhost_to_ipv4(base_url.rstrip("/"))
|
||||
if server_url.endswith("/v1"):
|
||||
server_url = server_url[:-3]
|
||||
server_url = _server_root(base_url)
|
||||
lmstudio_url = _localhost_to_ipv4(_lmstudio_server_root(base_url))
|
||||
|
||||
if _endpoint_blackholed(server_url):
|
||||
@@ -2418,32 +2343,15 @@ def _query_local_context_length_uncached(model: str, base_url: str, api_key: str
|
||||
|
||||
try:
|
||||
with httpx.Client(timeout=3.0, headers=headers) as client:
|
||||
# Ollama: /api/show returns model details with context info
|
||||
# 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().
|
||||
if server_type == "ollama":
|
||||
resp = client.post(f"{server_url}/api/show", json={"name": model})
|
||||
if resp.status_code == 200:
|
||||
data = resp.json()
|
||||
# Prefer explicit num_ctx from Modelfile parameters: this is
|
||||
# the *runtime* context Ollama will actually allocate KV cache
|
||||
# for. The GGUF model_info.context_length is the training max,
|
||||
# which can be larger than num_ctx — using it here would let
|
||||
# Hermes grow conversations past the runtime limit and Ollama
|
||||
# would silently truncate. Matches query_ollama_num_ctx().
|
||||
params = data.get("parameters", "")
|
||||
if "num_ctx" in params:
|
||||
for line in params.split("\n"):
|
||||
if "num_ctx" in line:
|
||||
parts = line.strip().split()
|
||||
if len(parts) >= 2:
|
||||
try:
|
||||
return int(parts[-1])
|
||||
except ValueError:
|
||||
pass
|
||||
# Fall back to GGUF model_info context_length (training max)
|
||||
model_info = data.get("model_info", {})
|
||||
for key, value in model_info.items():
|
||||
if "context_length" in key and isinstance(value, (int, float)):
|
||||
return int(value)
|
||||
ctx = _ollama_show_context(resp.json(), gguf_first=False)
|
||||
if ctx is not None:
|
||||
return ctx
|
||||
|
||||
# LM Studio native API: /api/v1/models returns max_context_length.
|
||||
# This is more reliable than the OpenAI-compat /v1/models which
|
||||
@@ -2879,19 +2787,6 @@ def _fetch_codex_oauth_context_lengths_with_source(
|
||||
return result, True
|
||||
|
||||
|
||||
def _fetch_codex_oauth_context_lengths(access_token: str) -> Dict[str, int]:
|
||||
"""Probe the ChatGPT Codex /models endpoint for per-slug context windows.
|
||||
|
||||
Codex OAuth imposes its own context limits that differ from the direct
|
||||
OpenAI API (e.g. gpt-5.5 is 1.05M on the API, 272K on Codex). The
|
||||
`context_window` field in each model entry is the authoritative source.
|
||||
|
||||
Returns a ``{slug: context_window}`` dict. Empty on failure.
|
||||
"""
|
||||
result, _fresh = _fetch_codex_oauth_context_lengths_with_source(access_token)
|
||||
return result
|
||||
|
||||
|
||||
def _resolve_codex_oauth_context_length_with_source(
|
||||
model: str, access_token: str = ""
|
||||
) -> Tuple[Optional[int], str]:
|
||||
@@ -2945,27 +2840,12 @@ def _resolve_codex_oauth_context_length_with_source(
|
||||
if slug.lower() == model_lower:
|
||||
return _apply_verified_bump(ctx, live_source)
|
||||
|
||||
# Fallback: longest-key-first substring match over hardcoded defaults.
|
||||
model_lower = lookup_bare.lower()
|
||||
for slug, ctx in sorted(
|
||||
_CODEX_OAUTH_CONTEXT_FALLBACK.items(), key=lambda x: len(x[0]), reverse=True
|
||||
):
|
||||
if slug in model_lower:
|
||||
return _apply_verified_bump(ctx, "fallback")
|
||||
|
||||
hit = _longest_key_match(_CODEX_OAUTH_CONTEXT_FALLBACK, lookup_bare.lower())
|
||||
if hit:
|
||||
return _apply_verified_bump(hit[1], "fallback")
|
||||
return None, ""
|
||||
|
||||
|
||||
def _resolve_codex_oauth_context_length(
|
||||
model: str, access_token: str = ""
|
||||
) -> Optional[int]:
|
||||
"""Resolve a Codex OAuth model's context length (compatibility wrapper)."""
|
||||
context_length, _source = _resolve_codex_oauth_context_length_with_source(
|
||||
model, access_token=access_token,
|
||||
)
|
||||
return context_length
|
||||
|
||||
|
||||
def _resolve_nous_context_length(
|
||||
model: str,
|
||||
base_url: str = "",
|
||||
@@ -3352,10 +3232,8 @@ def get_model_context_length(
|
||||
# would create a false-safe window for compression (#63122).
|
||||
# Non-local endpoints preserve the existing GGUF-first behavior.
|
||||
if is_local_endpoint(base_url):
|
||||
local_ctx = _query_local_context_length(model, base_url, api_key=api_key)
|
||||
if local_ctx and local_ctx > 0:
|
||||
if not _skip_persistent_context_cache(base_url, provider):
|
||||
_maybe_cache_local_context_length(model, base_url, local_ctx)
|
||||
local_ctx = _probe_local_context_length(model, base_url, api_key, provider)
|
||||
if local_ctx:
|
||||
return local_ctx
|
||||
# 2b. Ollama native /api/show — non-local endpoints preserve
|
||||
# the existing generic /api/show GGUF-first behavior.
|
||||
@@ -3379,19 +3257,14 @@ def get_model_context_length(
|
||||
# DEFAULT_CONTEXT_LENGTHS (e.g. "claude-opus-4-8" → 1M).
|
||||
# Without this, the early return here short-circuits the catalog
|
||||
# lookup at step 8 and silently caps context at 256K.
|
||||
model_lower = model.lower()
|
||||
for default_model, length in sorted(
|
||||
DEFAULT_CONTEXT_LENGTHS.items(),
|
||||
key=lambda x: len(x[0]),
|
||||
reverse=True,
|
||||
):
|
||||
if default_model in model_lower:
|
||||
logger.info(
|
||||
"Using hardcoded context length %s for model %r "
|
||||
"(custom endpoint, catalog match on %r)",
|
||||
f"{length:,}", model, default_model,
|
||||
)
|
||||
return length
|
||||
hit = _longest_key_match(DEFAULT_CONTEXT_LENGTHS, model.lower())
|
||||
if hit:
|
||||
logger.info(
|
||||
"Using hardcoded context length %s for model %r "
|
||||
"(custom endpoint, catalog match on %r)",
|
||||
f"{hit[1]:,}", model, hit[0],
|
||||
)
|
||||
return hit[1]
|
||||
# Same silent-256K bug class as the step-9 fallback below —
|
||||
# warn here too so custom/local endpoints aren't left invisible.
|
||||
_warn_context_length_fallback(model, base_url)
|
||||
@@ -3551,22 +3424,17 @@ def get_model_context_length(
|
||||
# ``Hermes-3-Llama-3.1-70B`` substring-match ``llama`` (131072) even when
|
||||
# vLLM is running at a lower ``--max-model-len`` (e.g. 32768 on limited VRAM).
|
||||
if base_url and is_local_endpoint(base_url):
|
||||
local_ctx = _query_local_context_length(model, base_url, api_key=api_key)
|
||||
if local_ctx and local_ctx > 0:
|
||||
if not _skip_persistent_context_cache(base_url, provider):
|
||||
_maybe_cache_local_context_length(model, base_url, local_ctx)
|
||||
local_ctx = _probe_local_context_length(model, base_url, api_key, provider)
|
||||
if local_ctx:
|
||||
return local_ctx
|
||||
|
||||
# 8. Hardcoded defaults (fuzzy match — longest key first for specificity)
|
||||
# Only check `default_model in model` (is the key a substring of the input).
|
||||
# The reverse (`model in default_model`) causes shorter names like
|
||||
# "claude-sonnet-4" to incorrectly match "claude-sonnet-4-6" and return 1M.
|
||||
model_lower = model.lower()
|
||||
for default_model, length in sorted(
|
||||
DEFAULT_CONTEXT_LENGTHS.items(), key=lambda x: len(x[0]), reverse=True
|
||||
):
|
||||
if default_model in model_lower:
|
||||
return length
|
||||
hit = _longest_key_match(DEFAULT_CONTEXT_LENGTHS, model.lower())
|
||||
if hit:
|
||||
return hit[1]
|
||||
|
||||
# 9. Default fallback — warn (deduped per model+endpoint) so
|
||||
# small-context models don't silently get 256K. See
|
||||
@@ -3911,17 +3779,6 @@ def _wire_message_shadow(msg: Dict[str, Any]) -> Dict[str, Any]:
|
||||
return shadow
|
||||
|
||||
|
||||
def _estimate_message_chars(msg: Dict[str, Any]) -> int:
|
||||
"""Char count for token estimation, excluding base64 image data.
|
||||
|
||||
Base64 images are counted via `_count_image_tokens` instead; including
|
||||
their raw chars here would massively overestimate token usage.
|
||||
"""
|
||||
if not isinstance(msg, dict):
|
||||
return len(str(msg))
|
||||
return len(str(_wire_message_shadow(msg)))
|
||||
|
||||
|
||||
def _estimate_message_tokens_without_images(msg: Dict[str, Any]) -> int:
|
||||
"""Token estimate for a message shadow with image payloads stripped."""
|
||||
if not isinstance(msg, dict):
|
||||
|
||||
Reference in New Issue
Block a user