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:
+72
-12
@@ -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
@@ -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
|
||||
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
|
||||
Reference in New Issue
Block a user