diff --git a/agent/model_metadata.py b/agent/model_metadata.py index 7acf3868c5..b1701c32e8 100644 --- a/agent/model_metadata.py +++ b/agent/model_metadata.py @@ -145,6 +145,96 @@ _ENDPOINT_MODEL_CACHE_TTL = 300 _ENDPOINT_PROBE_TTL_SECONDS = 3600.0 _endpoint_probe_path_cache: Dict[str, tuple] = {} +# A configured endpoint that is routable-but-dead — e.g. a corp LAN address +# while off-VPN — blackholes TCP: the SYN draws no SYN-ACK, no RST and no ICMP +# error, so a probe waits out its full timeout instead of failing fast. Startup +# runs a whole waterfall of such probes across several functions here, and the +# stalls stack into a minute-long hang before the banner renders. +# +# Once ANY probe has actually observed a connect timeout for an endpoint, the +# others have nothing to gain by repeating it. Recording that observation and +# short-circuiting on it performs no network I/O of its own — it adds no probe +# for callers or tests to mock, and it can only ever fire after a real timeout +# has already been paid, so it cannot suppress a probe that would have worked. +_ENDPOINT_BLACKHOLE_TTL_SECONDS = 30.0 +# Values are monotonic timestamps of the last observed connect timeout. +_endpoint_blackhole_cache: Dict[str, float] = {} + + +def _endpoint_host_key(base_url: str) -> Optional[str]: + """Return a ``host:port`` key for ``base_url``, or None if it has no host. + + Keyed on host:port rather than the full URL so every probe path for one + server — ``/v1``-suffixed or not, LM Studio root or API root — shares a + single entry. + """ + normalized = _normalize_base_url(base_url) + if not normalized: + return None + url = normalized if "://" in normalized else f"http://{normalized}" + try: + parsed = urlparse(url) + host = parsed.hostname + port = parsed.port or (443 if parsed.scheme == "https" else 80) + except Exception: + return None + return f"{host}:{port}" if host else None + + +def _note_endpoint_blackholed(base_url: str) -> None: + """Record that a probe to ``base_url`` timed out during TCP connect.""" + key = _endpoint_host_key(base_url) + if key is None: + return + _endpoint_blackhole_cache[key] = time.monotonic() + logger.debug( + "Endpoint %s timed out connecting — skipping further probes for %.0fs", + key, _ENDPOINT_BLACKHOLE_TTL_SECONDS, + ) + + +def _endpoint_blackholed(base_url: str) -> bool: + """True if a recent probe to ``base_url`` timed out during TCP connect. + + Pure cache lookup; never touches the network. The entry expires after + _ENDPOINT_BLACKHOLE_TTL_SECONDS — long enough to collapse one startup's + burst of probes, short enough that bringing the VPN up mid-session is + picked up without a restart. + """ + 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) + if seen is None: + return False + if (time.monotonic() - seen) >= _ENDPOINT_BLACKHOLE_TTL_SECONDS: + del _endpoint_blackhole_cache[key] + return False + return True + + +def _is_connect_timeout(exc: BaseException) -> bool: + """True for connect-phase timeouts raised by httpx or requests. + + Read timeouts are deliberately excluded: those mean the server accepted + the connection, which is the opposite of the blackhole this guards. + """ + try: + import httpx + if isinstance(exc, httpx.ConnectTimeout): + return True + except Exception: + pass + try: + from requests.exceptions import ConnectTimeout + if isinstance(exc, ConnectTimeout): + return True + except Exception: + pass + return False + # ── Disk L2 for local-endpoint probe results ──────────────────────────────── # The in-process caches above die with the process, so every CLI cold start # with a local model re-paid the probe waterfall in AIAgent.__init__: @@ -876,6 +966,13 @@ def detect_local_server_type(base_url: str, api_key: str = "") -> Optional[str]: if cached is not None and (time.monotonic() - cached[1]) < _ENDPOINT_PROBE_TTL_SECONDS: 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 + # written to _endpoint_probe_path_cache — that entry lives for an hour, + # which would pin the endpoint to "undetected" long after it comes back. + if _endpoint_blackholed(server_url): + return None + # Disk L2: a fresh cross-process verdict skips the HTTP waterfall # entirely (back-to-back CLI invocations, cron ticks). disk_hit = _local_probe_disk_get("server_type", server_url) @@ -885,6 +982,16 @@ def detect_local_server_type(base_url: str, api_key: str = "") -> Optional[str]: headers = _auth_headers(api_key) + def _probe_failed(exc: Exception) -> None: + """Swallow a probe error — or abort the waterfall if we were blackholed. + + Re-raising propagates out of the ``with`` block to the outer handler, + so the remaining legs are skipped instead of each stalling in turn. + """ + if _is_connect_timeout(exc): + _note_endpoint_blackholed(server_url) + raise exc + result: Optional[str] = None try: with httpx.Client(timeout=2.0, headers=headers) as client: @@ -893,8 +1000,8 @@ def detect_local_server_type(base_url: str, api_key: str = "") -> Optional[str]: r = client.get(f"{lmstudio_url}/api/v1/models") if r.status_code == 200: result = "lm-studio" - except Exception: - pass + 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 @@ -908,8 +1015,8 @@ def detect_local_server_type(base_url: str, api_key: str = "") -> Optional[str]: result = "ollama" except Exception: pass - 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: @@ -918,8 +1025,8 @@ def detect_local_server_type(base_url: str, api_key: str = "") -> Optional[str]: 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: - pass + except Exception as exc: + _probe_failed(exc) if result is None: # vLLM: /version try: @@ -928,8 +1035,8 @@ def detect_local_server_type(base_url: str, api_key: str = "") -> Optional[str]: data = r.json() if "version" in data: result = "vllm" - except Exception: - pass + except Exception as exc: + _probe_failed(exc) except Exception: pass @@ -1123,6 +1230,12 @@ def fetch_endpoint_model_metadata( if cached is not None and (time.time() - cached_at) < _ENDPOINT_MODEL_CACHE_TTL: return cached + # Blackholed endpoint: every candidate below would spend its full 5s + # connect budget. Returned empty rather than cached, so the endpoint is + # retried as soon as the blackhole entry expires. + if _endpoint_blackholed(normalized): + return {} + candidates = [normalized] if normalized.endswith("/v1"): alternate = normalized[:-3].rstrip("/") @@ -1185,8 +1298,15 @@ def fetch_endpoint_model_metadata( return cache except Exception as exc: last_error = exc + if _is_connect_timeout(exc): + _note_endpoint_blackholed(normalized) for candidate in candidates: + # A connect timeout on one candidate condemns the host, not the path: + # the remaining candidates differ only by URL suffix, so trying them + # would repeat the same stall. + if _endpoint_blackholed(normalized): + break # normalized/candidates stay unrewritten (cache key stability); only # the outbound request target is IPv4-resolved to skip the multi-second # dual-stack IPv6 connect timeout (see _localhost_to_ipv4). @@ -1257,6 +1377,8 @@ def fetch_endpoint_model_metadata( return cache except Exception as exc: last_error = exc + if _is_connect_timeout(exc): + _note_endpoint_blackholed(normalized) finally: if response is not None: response.close() @@ -1820,6 +1942,9 @@ def _query_ollama_api_show_uncached(model: str, base_url: str, api_key: str = "" if server_url.endswith("/v1"): server_url = server_url[:-3] + if _endpoint_blackholed(server_url): + return None + headers = _auth_headers(api_key) try: @@ -1851,8 +1976,9 @@ def _query_ollama_api_show_uncached(model: str, base_url: str, api_key: str = "" return ctx except ValueError: pass - except Exception: - pass + except Exception as exc: + if _is_connect_timeout(exc): + _note_endpoint_blackholed(server_url) return None @@ -1940,6 +2066,9 @@ def _query_local_context_length_uncached(model: str, base_url: str, api_key: str server_url = server_url[:-3] lmstudio_url = _localhost_to_ipv4(_lmstudio_server_root(base_url)) + if _endpoint_blackholed(server_url): + return None + headers = _auth_headers(api_key) try: @@ -2015,8 +2144,9 @@ def _query_local_context_length_uncached(model: str, base_url: str, api_key: str ctx = m.get("max_model_len") or m.get("context_length") or m.get("max_tokens") if ctx and isinstance(ctx, (int, float)): return int(ctx) - except Exception: - pass + except Exception as exc: + if _is_connect_timeout(exc): + _note_endpoint_blackholed(server_url) return None diff --git a/tests/agent/test_endpoint_blackhole.py b/tests/agent/test_endpoint_blackhole.py new file mode 100644 index 0000000000..b3bdc93a08 --- /dev/null +++ b/tests/agent/test_endpoint_blackhole.py @@ -0,0 +1,291 @@ +"""Tests for short-circuiting probes to endpoints that blackhole TCP connects. + +A routable-but-dead endpoint (e.g. a corp LAN address while off-VPN) drops SYNs +without a RST or ICMP error, so each probe waits out its full timeout. Once one +probe has observed that, the rest must not repeat it. + +Covers: +- _endpoint_blackholed / _note_endpoint_blackholed host:port keying and TTL +- detect_local_server_type aborting its waterfall on the first connect timeout +- fetch_endpoint_model_metadata skipping its candidate loop once blackholed +- _query_ollama_api_show_uncached / _query_local_context_length_uncached + honouring and recording the blackhole +- non-timeout failures (refused, no route) leaving the waterfall untouched +""" + +from __future__ import annotations + +import os +import sys +from unittest.mock import MagicMock, patch + +import httpx +import pytest +import requests + +sys.path.insert(0, os.path.join(os.path.dirname(__file__), "..")) + + +@pytest.fixture(autouse=True) +def _clear_caches(): + """Module-level caches must not leak between tests.""" + from agent import model_metadata + model_metadata._endpoint_blackhole_cache.clear() + model_metadata._endpoint_probe_path_cache.clear() + model_metadata._endpoint_model_metadata_cache.clear() + model_metadata._endpoint_model_metadata_cache_time.clear() + model_metadata._LOCAL_CTX_PROBE_CACHE.clear() + yield + model_metadata._endpoint_blackhole_cache.clear() + model_metadata._endpoint_probe_path_cache.clear() + model_metadata._endpoint_model_metadata_cache.clear() + model_metadata._endpoint_model_metadata_cache_time.clear() + model_metadata._LOCAL_CTX_PROBE_CACHE.clear() + + +def _client_mock(side_effect): + client = MagicMock() + client.__enter__ = lambda s: client + client.__exit__ = MagicMock(return_value=False) + client.get.side_effect = side_effect + client.post.side_effect = side_effect + return client + + +class TestBlackholeCache: + def test_unseen_endpoint_is_not_blackholed(self): + from agent.model_metadata import _endpoint_blackholed + + assert _endpoint_blackholed("http://10.0.0.9:30080/v1") is False + + def test_note_then_detected(self): + from agent.model_metadata import _endpoint_blackholed, _note_endpoint_blackholed + + _note_endpoint_blackholed("http://10.0.0.9:30080/v1") + assert _endpoint_blackholed("http://10.0.0.9:30080/v1") is True + + def test_keyed_on_host_port_not_path(self): + """Every probe path for one server shares a single entry.""" + from agent.model_metadata import _endpoint_blackholed, _note_endpoint_blackholed + + _note_endpoint_blackholed("http://10.0.0.9:30080") + assert _endpoint_blackholed("http://10.0.0.9:30080/v1") is True + assert _endpoint_blackholed("http://10.0.0.9:30080/api/v1") is True + + def test_different_port_is_independent(self): + from agent.model_metadata import _endpoint_blackholed, _note_endpoint_blackholed + + _note_endpoint_blackholed("http://10.0.0.9:30080/v1") + assert _endpoint_blackholed("http://10.0.0.9:11434/v1") is False + + def test_entry_expires_after_ttl(self): + """A recovered endpoint (VPN back up) is probed again without a restart.""" + from agent import model_metadata + from agent.model_metadata import _endpoint_blackholed, _note_endpoint_blackholed + + _note_endpoint_blackholed("http://10.0.0.9:30080/v1") + stale = ( + model_metadata._endpoint_blackhole_cache["10.0.0.9:30080"] + - model_metadata._ENDPOINT_BLACKHOLE_TTL_SECONDS + - 1 + ) + model_metadata._endpoint_blackhole_cache["10.0.0.9:30080"] = stale + assert _endpoint_blackholed("http://10.0.0.9:30080/v1") is False + + def test_ttl_zero_disables_short_circuit(self): + from agent import model_metadata + from agent.model_metadata import _endpoint_blackholed, _note_endpoint_blackholed + + _note_endpoint_blackholed("http://10.0.0.9:30080/v1") + with patch.object(model_metadata, "_ENDPOINT_BLACKHOLE_TTL_SECONDS", 0.0): + assert _endpoint_blackholed("http://10.0.0.9:30080/v1") is False + + +class TestDetectLocalServerTypeBlackhole: + URL = "http://10.0.0.9:30080/v1" + + def test_connect_timeout_aborts_waterfall_after_one_probe(self): + """Four sequential 2s probes against a dead host must collapse to one.""" + from agent.model_metadata import _endpoint_blackholed, detect_local_server_type + + client = _client_mock(httpx.ConnectTimeout("timed out")) + with patch("httpx.Client", return_value=client): + assert detect_local_server_type(self.URL) is None + + assert client.get.call_count == 1 + assert _endpoint_blackholed(self.URL) is True + + def test_second_call_makes_no_request_at_all(self): + from agent.model_metadata import detect_local_server_type + + client = _client_mock(httpx.ConnectTimeout("timed out")) + with patch("httpx.Client", return_value=client): + detect_local_server_type(self.URL) + first_count = client.get.call_count + assert detect_local_server_type(self.URL) is None + + assert client.get.call_count == first_count + + def test_refused_does_not_blackhole_and_runs_full_waterfall(self): + """Refused answers instantly, so skipping buys nothing and must not fire. + + This is the common "local server not started yet" path. + """ + from agent.model_metadata import _endpoint_blackholed, detect_local_server_type + + client = _client_mock(httpx.ConnectError("connection refused")) + with patch("httpx.Client", return_value=client): + assert detect_local_server_type(self.URL) is None + + assert client.get.call_count > 1 + assert _endpoint_blackholed(self.URL) is False + + def test_read_timeout_does_not_blackhole(self): + """A read timeout means the connection was accepted — not a blackhole.""" + from agent.model_metadata import _endpoint_blackholed, detect_local_server_type + + client = _client_mock(httpx.ReadTimeout("slow")) + with patch("httpx.Client", return_value=client): + detect_local_server_type(self.URL) + + assert _endpoint_blackholed(self.URL) is False + + +class TestFetchEndpointModelMetadataBlackhole: + URL = "http://10.0.0.9:30080/v1" + + def test_connect_timeout_skips_remaining_candidates(self): + """A timeout condemns the host, not the URL suffix — one stall, not two.""" + from agent.model_metadata import _endpoint_blackholed, fetch_endpoint_model_metadata + + with patch("agent.model_metadata.detect_local_server_type", return_value=None), \ + patch( + "agent.model_metadata.requests.get", + side_effect=requests.exceptions.ConnectTimeout("timed out"), + ) as get: + assert fetch_endpoint_model_metadata(self.URL) == {} + + assert get.call_count == 1 + assert _endpoint_blackholed(self.URL) is True + + def test_refused_tries_every_candidate_and_does_not_blackhole(self): + from agent.model_metadata import _endpoint_blackholed, fetch_endpoint_model_metadata + + with patch("agent.model_metadata.detect_local_server_type", return_value=None), \ + patch( + "agent.model_metadata.requests.get", + side_effect=requests.exceptions.ConnectionError("refused"), + ) as get: + assert fetch_endpoint_model_metadata(self.URL) == {} + + assert get.call_count == 2 # /v1-suffixed and bare candidates + assert _endpoint_blackholed(self.URL) is False + + def test_blackholed_endpoint_issues_no_request(self): + """force_refresh bypasses the metadata cache, so only the guard can stop it.""" + from agent.model_metadata import _note_endpoint_blackholed, fetch_endpoint_model_metadata + + _note_endpoint_blackholed(self.URL) + with patch("agent.model_metadata.detect_local_server_type", return_value=None), \ + patch("agent.model_metadata.requests.get") as get: + assert fetch_endpoint_model_metadata(self.URL, force_refresh=True) == {} + + get.assert_not_called() + + +class TestQueryOllamaApiShowBlackhole: + URL = "http://10.0.0.9:30080/v1" + + def test_connect_timeout_records_blackhole(self): + from agent.model_metadata import _endpoint_blackholed, _query_ollama_api_show_uncached + + client = _client_mock(httpx.ConnectTimeout("timed out")) + with patch("httpx.Client", return_value=client): + assert _query_ollama_api_show_uncached("some-model", self.URL) is None + + assert client.post.call_count == 1 + assert _endpoint_blackholed(self.URL) is True + + def test_blackholed_endpoint_issues_no_request(self): + from agent.model_metadata import _note_endpoint_blackholed, _query_ollama_api_show_uncached + + _note_endpoint_blackholed(self.URL) + with patch("httpx.Client") as client_cls: + assert _query_ollama_api_show_uncached("some-model", self.URL) is None + + client_cls.assert_not_called() + + def test_read_timeout_does_not_blackhole(self): + from agent.model_metadata import _endpoint_blackholed, _query_ollama_api_show_uncached + + client = _client_mock(httpx.ReadTimeout("slow")) + with patch("httpx.Client", return_value=client): + assert _query_ollama_api_show_uncached("some-model", self.URL) is None + + assert _endpoint_blackholed(self.URL) is False + + +class TestQueryLocalContextLengthBlackhole: + URL = "http://10.0.0.9:30080/v1" + + def test_connect_timeout_records_blackhole(self): + from agent.model_metadata import ( + _endpoint_blackholed, + _query_local_context_length_uncached, + ) + + client = _client_mock(httpx.ConnectTimeout("timed out")) + with patch("agent.model_metadata.detect_local_server_type", return_value=None), \ + patch("httpx.Client", return_value=client): + assert _query_local_context_length_uncached("some-model", self.URL) is None + + assert _endpoint_blackholed(self.URL) is True + + def test_blackholed_endpoint_skips_detection_and_requests(self): + """The guard sits before detect_local_server_type — nothing runs at all.""" + from agent.model_metadata import ( + _note_endpoint_blackholed, + _query_local_context_length_uncached, + ) + + _note_endpoint_blackholed(self.URL) + with patch("agent.model_metadata.detect_local_server_type") as detect, \ + patch("httpx.Client") as client_cls: + assert _query_local_context_length_uncached("some-model", self.URL) is None + + detect.assert_not_called() + client_cls.assert_not_called() + + def test_read_timeout_does_not_blackhole(self): + from agent.model_metadata import ( + _endpoint_blackholed, + _query_local_context_length_uncached, + ) + + client = _client_mock(httpx.ReadTimeout("slow")) + with patch("agent.model_metadata.detect_local_server_type", return_value=None), \ + patch("httpx.Client", return_value=client): + assert _query_local_context_length_uncached("some-model", self.URL) is None + + assert _endpoint_blackholed(self.URL) is False + + +class TestIsConnectTimeout: + def test_httpx_connect_timeout(self): + from agent.model_metadata import _is_connect_timeout + + assert _is_connect_timeout(httpx.ConnectTimeout("x")) is True + + def test_requests_connect_timeout(self): + from requests.exceptions import ConnectTimeout + + from agent.model_metadata import _is_connect_timeout + + assert _is_connect_timeout(ConnectTimeout("x")) is True + + def test_unrelated_errors_are_not_connect_timeouts(self): + from agent.model_metadata import _is_connect_timeout + + assert _is_connect_timeout(httpx.ReadTimeout("x")) is False + assert _is_connect_timeout(httpx.ConnectError("x")) is False + assert _is_connect_timeout(ValueError("x")) is False