fix(vision): forward the API key to the server-type probe and cache failed verdicts

The image-routing vision path calls detect_local_server_type without
the provider's API key. Against a remote API-keyed endpoint (sglang /
vLLM with --api-key) every leg of the 5-request probe waterfall came
back 401 — and because a failed verdict was never written to the
in-memory cache (only positive verdicts were), the waterfall re-ran on
EVERY image-bearing turn (#89863: 51 detail-less busy-acks observed in
one Slack channel while the probe sprayed the user's own server).

Two changes:

- image_routing._should_probe_ollama_vision now takes the API key and
  forwards it; a new _resolve_inference_api_key mirrors
  _resolve_inference_base_url's resolution order (runtime value,
  model.api_key, providers blocks) so the key always matches the URL
  being probed.

- detect_local_server_type caches a None verdict in memory with a short
  failure TTL (5 min, vs 1h for positives) so the next turn is served
  from the negative entry instead of re-running the waterfall — while
  a transient failure (server starting, key being fixed) recovers in
  minutes. Negative verdicts are deliberately not written to the
  cross-process disk cache.
This commit is contained in:
liuhao1024
2026-08-19 16:25:27 +08:00
committed by Teknium
parent 4ca993c746
commit 51239e8e2a
4 changed files with 241 additions and 14 deletions
+72 -12
View File
@@ -332,10 +332,73 @@ def _resolve_inference_base_url(
return ""
def _should_probe_ollama_vision(
def _resolve_inference_api_key(
cfg: Optional[Dict[str, Any]],
provider: str,
base_url: str,
api_key: str = "",
) -> str:
"""Best-effort API key for the active inference provider.
Mirrors :func:`_resolve_inference_base_url`'s resolution order (runtime
value, then ``model.api_key``, then the providers blocks) so the key
matches the base URL actually being probed. Without this, the local
server-type probe fires at a remote API-keyed endpoint without an
Authorization header — 5×401 per image-bearing turn on a keyed
sglang/vLLM deployment (#89863).
"""
try:
from agent.auxiliary_client import _runtime_main_value
runtime_key = str(_runtime_main_value("api_key") or "").strip()
if runtime_key:
return runtime_key
except Exception:
pass
if not isinstance(cfg, dict):
return ""
model_cfg_raw = cfg.get("model")
model_cfg: Dict[str, Any] = model_cfg_raw if isinstance(model_cfg_raw, dict) else {}
key = str(model_cfg.get("api_key") or "").strip()
if key:
return key
config_provider = str(model_cfg.get("provider") or "").strip()
candidate_names: set[str] = set()
for p in filter(None, (provider, config_provider)):
candidate_names.add(p)
if p.lower().startswith("custom:"):
candidate_names.add(p.split(":", 1)[1])
else:
candidate_names.add(f"custom:{p}")
providers_cfg = cfg.get("providers")
if isinstance(providers_cfg, dict):
for name in candidate_names:
entry = providers_cfg.get(name)
if isinstance(entry, dict):
k = str(entry.get("api_key") or "").strip()
if k:
return k
custom_providers = cfg.get("custom_providers")
if isinstance(custom_providers, list):
lowered = {n.lower() for n in candidate_names}
for entry_raw in custom_providers:
if not isinstance(entry_raw, dict):
continue
entry_name = str(entry_raw.get("name") or "").strip()
if entry_name not in candidate_names and entry_name.lower() not in lowered:
continue
k = str(entry_raw.get("api_key") or "").strip()
if k:
return k
return ""
def _should_probe_ollama_vision(
provider: str, base_url: str, api_key: str = ""
) -> bool:
"""True when the active provider likely fronts a local Ollama server.
@@ -364,6 +427,9 @@ def _should_probe_ollama_vision(
try:
from agent.model_metadata import detect_local_server_type
# Forward the API key: a remote API-keyed endpoint answers the
# probe waterfall with 401s without it, and an unauthorized probe
# can never produce a positive verdict (#89863).
return detect_local_server_type(base_url, api_key=api_key) == "ollama"
except Exception:
return False
@@ -471,15 +537,9 @@ def _lookup_supports_vision(
if not base_url and (provider or "").strip().lower() == "ollama":
base_url = "http://localhost:11434/v1"
# Resolve the runtime api_key so probe requests at keyed endpoints carry
# Authorization and don't spray 401s (issue #89863).
resolved_api_key = ""
try:
from agent.auxiliary_client import _runtime_main_value
resolved_api_key = str(_runtime_main_value("api_key") or "").strip()
except Exception:
pass
# Resolve the provider's API key so probe requests at keyed endpoints
# carry Authorization and don't spray 401s (issue #89863).
resolved_api_key = _resolve_inference_api_key(cfg, provider)
if _should_probe_ollama_vision(provider, base_url, api_key=resolved_api_key):
try:
+26 -2
View File
@@ -158,6 +158,14 @@ _ENDPOINT_MODEL_CACHE_TTL = 300
# of being pinned to the stale type for the whole process lifetime.
# Values are (server_type, monotonic_timestamp).
_ENDPOINT_PROBE_TTL_SECONDS = 3600.0
# A failed probe verdict (server_type is None — no known endpoint answered)
# is cached for a much shorter window: the in-memory entry exists only to
# keep one image-bearing turn from re-running the 5-request waterfall on
# every subsequent turn (#89863 — a keyed remote endpoint answered 401 to
# each leg and the None verdict was never cached, so every turn re-probed).
# Short TTL keeps a transient failure (server starting up, key being fixed)
# recoverable within minutes instead of pinning "undetected" for an hour.
_ENDPOINT_PROBE_FAILURE_TTL_SECONDS = 300.0
_endpoint_probe_path_cache: Dict[str, tuple] = {}
# A configured endpoint that is routable-but-dead — e.g. a corp LAN address
@@ -1009,8 +1017,18 @@ def detect_local_server_type(base_url: str, api_key: str = "") -> Optional[str]:
lmstudio_url = _lmstudio_server_root(normalized)
cached = _endpoint_probe_path_cache.get(server_url)
if cached is not None and (time.monotonic() - cached[1]) < _ENDPOINT_PROBE_TTL_SECONDS:
return cached[0]
if cached is not None:
# Positive verdicts live for the full TTL; a None verdict (probe
# waterfall answered nothing recognizable) gets the short failure
# TTL so it still throttles re-probing without pinning the
# endpoint as undetected for a whole hour (#89863).
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]
# The host already blackholed a connect: skip the waterfall below, each leg
# of which would otherwise burn its full 2s timeout. Deliberately NOT
@@ -1089,6 +1107,12 @@ def detect_local_server_type(base_url: str, api_key: str = "") -> Optional[str]:
if result is not None:
_endpoint_probe_path_cache[server_url] = (result, time.monotonic())
_local_probe_disk_put("server_type", server_url, result)
else:
# Cache the negative verdict in memory only (never on disk — a
# failure is often transient: server starting, key being fixed)
# so the very next turn does not re-run the whole waterfall
# against an endpoint that just answered nothing (#89863).
_endpoint_probe_path_cache[server_url] = (None, time.monotonic())
return result
+80
View File
@@ -571,3 +571,83 @@ class TestCustomProviderVisionAlias:
]
}
assert _supports_vision_override(cfg, "custom:my-vllm", "other") is None
def _fake_key(tag: str) -> str:
"""Build an obviously-fake placeholder key from parts — never a literal
that could be mistaken for (or collide with) a real credential."""
return "fake-" + tag + "-not-a-secret"
class TestProbeApiKeyForwarding:
"""The local server-type probe must carry the provider's API key (#89863).
A remote API-keyed endpoint answers the probe waterfall with 401s
without the Authorization header, and an unauthorized probe can never
produce a positive verdict — so every image-bearing turn re-ran the
5-request waterfall against the user's own server.
"""
def test_resolve_inference_api_key_model_block(self):
from agent.image_routing import _resolve_inference_api_key
key = _fake_key("model")
cfg = {"model": {"api_key": key}}
assert _resolve_inference_api_key(cfg, "custom") == key
def test_resolve_inference_api_key_providers_block(self):
from agent.image_routing import _resolve_inference_api_key
key = _fake_key("prov")
cfg = {
"model": {"provider": "custom:remote"},
"providers": {
"custom:remote": {"base_url": "https://x/v1", "api_key": key}
},
}
assert _resolve_inference_api_key(cfg, "custom:remote") == key
def test_resolve_inference_api_key_custom_providers_list(self):
from agent.image_routing import _resolve_inference_api_key
key = _fake_key("list")
cfg = {
"model": {"provider": "remote"},
"custom_providers": [
{"name": "remote", "base_url": "https://x/v1", "api_key": key}
],
}
assert _resolve_inference_api_key(cfg, "remote") == key
def test_resolve_inference_api_key_absent(self):
from agent.image_routing import _resolve_inference_api_key
assert _resolve_inference_api_key({"model": {}}, "custom") == ""
assert _resolve_inference_api_key(None, "custom") == ""
def test_should_probe_forwards_api_key(self):
from agent.image_routing import _should_probe_ollama_vision
key = _fake_key("probe")
with patch(
"agent.model_metadata.detect_local_server_type",
return_value=None,
) as detect:
_should_probe_ollama_vision("custom", "https://remote/v1", api_key=key)
detect.assert_called_once_with("https://remote/v1", api_key=key)
def test_lookup_passes_resolved_key_to_probe(self):
"""The full lookup path resolves the key from cfg and hands it to the
probe — the exact chain that sprayed 401s in #89863."""
key = _fake_key("lookup")
import agent.models_dev # noqa: F401 — make the patch target importable
with patch(
"agent.models_dev.get_model_capabilities", return_value=None
), patch(
"agent.image_routing._resolve_inference_base_url",
return_value="https://remote/v1",
), patch(
"agent.model_metadata.detect_local_server_type", return_value=None
) as detect:
_lookup_supports_vision("custom", "llava", {"model": {"api_key": key}})
assert detect.call_args.kwargs.get("api_key") == key
+63
View File
@@ -279,3 +279,66 @@ class TestContextCacheKeyNormalization:
assert "m1@http://host/v1/" not in cache
class TestDetectServerTypeNegativeCaching:
"""A failed detect_local_server_type verdict is cached briefly (#89863).
Previously only positive verdicts were memoized, so a remote endpoint
that answered the whole waterfall with 401s (no recognizable server
type) was re-probed — 5 requests — on every image-bearing turn.
"""
@staticmethod
def _client_all_401():
client = MagicMock()
client.__enter__ = lambda s: client
client.__exit__ = MagicMock(return_value=False)
resp = MagicMock()
resp.status_code = 401
client.get.return_value = resp
return client
def test_negative_verdict_is_cached_in_memory(self):
from agent.model_metadata import detect_local_server_type
from agent import model_metadata
client = self._client_all_401()
with patch("httpx.Client", return_value=client):
assert detect_local_server_type("http://remote:8080/v1") is None
assert detect_local_server_type("http://remote:8080/v1") is None
# Second call served from the in-memory negative entry: the
# waterfall ran exactly once (5 GETs), not twice.
assert client.get.call_count == 5
assert "http://remote:8080" in model_metadata._endpoint_probe_path_cache
def test_negative_verdict_not_written_to_disk(self):
from agent.model_metadata import detect_local_server_type
from agent import model_metadata
with patch("httpx.Client", return_value=self._client_all_401()), patch.object(
model_metadata, "_local_probe_disk_put"
) as disk_put:
assert detect_local_server_type("http://remote2:8080/v1") is None
disk_put.assert_not_called()
def test_negative_verdict_expires_quickly(self):
"""The short failure TTL keeps a transient failure recoverable."""
import time as _time
from agent.model_metadata import detect_local_server_type
from agent import model_metadata
client = self._client_all_401()
with patch("httpx.Client", return_value=client):
assert detect_local_server_type("http://remote3:8080/v1") is None
# Age the entry past the failure TTL.
model_metadata._endpoint_probe_path_cache["http://remote3:8080"] = (
None,
_time.monotonic()
- model_metadata._ENDPOINT_PROBE_FAILURE_TTL_SECONDS
- 1,
)
assert detect_local_server_type("http://remote3:8080/v1") is None
assert client.get.call_count == 10 # waterfall re-ran after expiry