Address review: endpoint-aware DNS check, shorter memo, strict gateway mode

Four issues raised on the first commit, all reproduced before fixing:

- The DNS pre-check hardcoded huggingface.co, so a reachable HF_ENDPOINT
  mirror was forced offline whenever huggingface.co did not resolve. It now
  follows the configured endpoint.
- The reachability verdict was memoised for 60s, and a stale "reachable"
  hid the user pulling the plug right after a download, which is the exact
  workflow this fix targets. Window is now 5s in both directions: long
  enough to dedupe the probes within one load, short enough that neither
  direction goes stale.
- hf_endpoint_unreachable counts 502/503/504 as offline, and the training
  worker used it to set flags for the whole job, so a momentary hub blip
  blocked every download for the rest of the run. Added
  gateway_errors_offline=False for callers setting lifetime flags; scoped
  callers keep the existing behaviour.
- Dropped the guard from _target_is_vision. The resolver only yields local
  paths there, so it returns from the mmproj filesystem branch without
  touching the hub, and the probe only added latency per request.

Verified unchanged offline: load 686s -> 5s, /models/config 378s -> 0s,
/models/check-vision 28s -> 0s.
This commit is contained in:
Unsloth 2026-07-28 22:30:24 -07:00
commit c4973c76f6
6 changed files with 142 additions and 14 deletions

View file

@ -495,9 +495,23 @@ _SWA_CACHE: Optional[dict] = None
_SWA_CACHE_LOCK = threading.Lock()
def _probe_dns_dead(host: str = "huggingface.co", timeout: float = 2.0) -> bool:
def _hf_endpoint_host() -> str:
"""Host of the configured hub endpoint. Mirror users point HF_ENDPOINT elsewhere, and
probing huggingface.co would then report their working mirror as offline."""
endpoint = (os.environ.get("HF_ENDPOINT") or "").strip() or "https://huggingface.co"
if "://" not in endpoint:
endpoint = "https://" + endpoint
try:
from urllib.parse import urlparse
return urlparse(endpoint).hostname or "huggingface.co"
except Exception:
return "huggingface.co"
def _probe_dns_dead(host: Optional[str] = None, timeout: float = 2.0) -> bool:
"""Quick DNS check on a daemon thread, so concurrent sockets aren't
affected by socket.setdefaulttimeout."""
affected by socket.setdefaulttimeout. Defaults to the configured endpoint's host."""
host = host or _hf_endpoint_host()
result: list[Optional[bool]] = [None]
def _probe() -> None:

View file

@ -2365,10 +2365,14 @@ def run_training_process(*, event_queue: Any, stop_queue: Any, config: dict) ->
if _result[0] is False:
# DNS answers even when there is no egress (WAN down, captive portal), so
# confirm with the bounded, proxy-aware reachability probe. HF_ENDPOINT aware.
# These flags last the whole job, so only a connection failure counts: a
# momentary 502/503 must not block every download for the rest of the run.
try:
from utils.transformers_version import hf_endpoint_unreachable
from utils.utils import hf_probe_disabled
if not hf_probe_disabled() and hf_endpoint_unreachable():
if not hf_probe_disabled() and hf_endpoint_unreachable(
gateway_errors_offline = False
):
_result[0] = True
except Exception:
pass

View file

@ -3672,10 +3672,10 @@ def _target_is_vision(load_path: str) -> bool:
# paths, where the token is unused, but the rule requires it regardless).
from utils.models.model_config import is_vision_model
try:
# Guarded: this runs per request, so an unreachable hub would re-pay its retry
# backoff on every image/audio call.
with _hf_offline_if_unreachable():
return bool(is_vision_model(load_path, hf_token = os.environ.get("HF_TOKEN")))
# Deliberately unguarded: the resolver only yields local paths, so this returns
# from the mmproj filesystem branch without touching the hub. A reachability
# probe here would add seconds per request and prevent nothing.
return bool(is_vision_model(load_path, hf_token = os.environ.get("HF_TOKEN")))
except Exception as exc:
# Detection failure: don't block the swap, let the load decide.
logger.debug("auto-switch: vision probe failed for %s: %s", load_path, exc)

View file

@ -1012,6 +1012,80 @@ class TestHfOfflineIfUnreachable:
assert "TRANSFORMERS_OFFLINE" not in os.environ
class TestEndpointAwareOfflineDetection:
"""A reachable HF_ENDPOINT mirror must not be declared offline just because
huggingface.co does not resolve (air-gapped / corporate networks)."""
@pytest.fixture
def no_upstream_dns(self, monkeypatch):
real_host = socket.gethostbyname
def _host(h, *a, **k):
if "huggingface.co" in str(h):
raise socket.gaierror(-2, "Name or service not known")
return "127.0.0.1"
monkeypatch.setattr(socket, "gethostbyname", _host)
@pytest.mark.parametrize(
"endpoint,expected",
[
("https://hf-mirror.com", "hf-mirror.com"),
("hf-mirror.com", "hf-mirror.com"),
("https://hf-mirror.com:8443", "hf-mirror.com"),
("https://hf-mirror.com/path", "hf-mirror.com"),
("", "huggingface.co"),
],
)
def test_endpoint_host_parsing(self, monkeypatch, endpoint, expected):
from core.inference.llama_cpp import _hf_endpoint_host
monkeypatch.setenv("HF_ENDPOINT", endpoint)
assert _hf_endpoint_host() == expected
def test_dns_precheck_follows_endpoint(self, monkeypatch, no_upstream_dns):
monkeypatch.setenv("HF_ENDPOINT", "https://hf-mirror.com")
assert _probe_dns_dead() is False
def test_default_endpoint_still_probes_huggingface(self, monkeypatch, no_upstream_dns):
monkeypatch.delenv("HF_ENDPOINT", raising = False)
assert _probe_dns_dead() is True
class TestGatewayErrorsAreNotConnectionFailures:
"""Lifetime offline flags must not be set by a momentary 502/503/504."""
def _probe_with(self, monkeypatch, exc):
import urllib.request
def _urlopen(*a, **k):
raise exc
monkeypatch.setattr(urllib.request, "urlopen", _urlopen)
from utils.transformers_version import hf_endpoint_unreachable
return hf_endpoint_unreachable
@pytest.mark.parametrize("code", [502, 503, 504])
def test_strict_mode_treats_gateway_error_as_reachable(self, monkeypatch, code):
import urllib.error
exc = urllib.error.HTTPError("u", code, "err", {}, None)
probe = self._probe_with(monkeypatch, exc)
assert probe(timeout = 1, gateway_errors_offline = False) is False
# Default (scoped callers) keeps treating a downed hub as offline.
assert probe(timeout = 1) is True
@pytest.mark.parametrize("code", [401, 403, 404, 429])
def test_other_http_errors_always_reachable(self, monkeypatch, code):
import urllib.error
exc = urllib.error.HTTPError("u", code, "err", {}, None)
probe = self._probe_with(monkeypatch, exc)
assert probe(timeout = 1) is False
assert probe(timeout = 1, gateway_errors_offline = False) is False
class TestHfUnreachableProbe:
"""``utils.utils.hf_unreachable``: memoised, opt-outable, fails open."""
@ -1069,6 +1143,31 @@ class TestHfUnreachableProbe:
assert hf_unreachable() is True
assert len(calls) == 2
def test_memo_window_is_short_in_both_directions(self):
"""Stale either way is a bug: a stale 'reachable' hides the plug being pulled,
a stale 'unreachable' fails a download after the user reconnects."""
import utils.utils as uu
assert uu._HF_REACHABILITY_TTL_S <= 10.0
def test_verdict_expires_so_a_disconnect_is_noticed(self, monkeypatch, clean_offline_env):
import time as _time
import utils.utils as uu
from utils.utils import hf_unreachable
monkeypatch.setattr(uu, "_HF_REACHABILITY_TTL_S", 0.2)
verdict = {"value": False}
monkeypatch.setattr(
__import__("utils.transformers_version", fromlist = ["x"]),
"hf_endpoint_unreachable",
lambda *a, **k: verdict["value"],
)
assert hf_unreachable() is False # online during the download
verdict["value"] = True # plug pulled
_time.sleep(0.3)
assert hf_unreachable() is True
class TestExtractQuantLabelSubdir:
"""``_extract_quant_label`` must consider parent dirs when the basename has

View file

@ -63,7 +63,7 @@ def _env_offline() -> bool:
)
def hf_endpoint_unreachable(timeout: int = 3) -> bool:
def hf_endpoint_unreachable(timeout: int = 3, *, gateway_errors_offline: bool = True) -> bool:
"""Bounded reachability probe to the HF endpoint. A HEAD request runs in a daemon thread
joined with a deadline, so a resolver blackhole cannot block past ~timeout+1s. True if
unreachable. urllib natively honors *_PROXY / NO_PROXY, so this verifies real egress
@ -86,8 +86,11 @@ def hf_endpoint_unreachable(timeout: int = 3) -> bool:
with urllib.request.urlopen(req, timeout = timeout):
result["online"] = True
except urllib.error.HTTPError as exc:
# The server/proxy answered: reachable unless it is a gateway error.
result["online"] = exc.code not in (502, 503, 504)
# The server/proxy answered, so we have egress. A gateway error usually means
# the hub itself is down, which callers scoping offline to one operation want
# to treat as offline; callers setting a lifetime flag pass
# gateway_errors_offline=False so a momentary 503 can't strand the process.
result["online"] = True if not gateway_errors_offline else exc.code not in (502, 503, 504)
except urllib.error.URLError as exc:
# A TLS/cert failure means we DID reach the server; treat as reachable so the real
# load surfaces it (consistent with _is_offline_related_error not retrying TLS).

View file

@ -37,12 +37,20 @@ def hf_env_offline() -> bool:
return False
# One load makes many hub calls, so the reachability verdict is shared for a short window.
_HF_REACHABILITY_TTL_S = 60.0
# One load makes many hub calls, so the verdict is shared briefly to avoid re-probing on
# each. Kept short in BOTH directions: a stale "reachable" misses the plug being pulled
# (the case this whole path exists for), and a stale "unreachable" sends a load to the
# cache after the user reconnected, failing it if the model is not cached.
_HF_REACHABILITY_TTL_S = 5.0
_hf_reachability: Optional[tuple] = None
_hf_reachability_lock = threading.Lock()
def _reachability_fresh(entry) -> bool:
"""True while a cached (timestamp, unreachable) verdict may still be reused."""
return entry is not None and (time.monotonic() - entry[0]) < _HF_REACHABILITY_TTL_S
def hf_probe_disabled() -> bool:
"""True when UNSLOTH_OFFLINE_PROBE opts out of the reachability probe."""
return os.environ.get("UNSLOTH_OFFLINE_PROBE", "1").strip().lower() in {
@ -75,12 +83,12 @@ def hf_unreachable(timeout: int = 3) -> bool:
global _hf_reachability
cached = _hf_reachability
if cached is not None and time.monotonic() - cached[0] < _HF_REACHABILITY_TTL_S:
if _reachability_fresh(cached):
return cached[1]
with _hf_reachability_lock:
cached = _hf_reachability
if cached is not None and time.monotonic() - cached[0] < _HF_REACHABILITY_TTL_S:
if _reachability_fresh(cached):
return cached[1]
try:
from utils.transformers_version import hf_endpoint_unreachable