fix(offline-guard): raise instead of deadlocking when a thread nests opposite modes

This commit is contained in:
Roman Dolgov
2026-10-03 09:18:41 +00:00
committed by jamiepine
parent e3b5e9b258
commit 291a4d8b97
2 changed files with 163 additions and 97 deletions
+31
View File
@@ -155,5 +155,36 @@ def test_uncached_load_never_observes_offline_flag_from_concurrent_cached_load()
assert original == _hf_const().HF_HUB_OFFLINE
def test_nesting_opposite_mode_on_same_thread_raises_instead_of_deadlocking():
"""Nesting is_cached=True inside is_cached=False (or vice versa) on the
same thread must raise immediately, not hang: the inner call's wait
condition can only be cleared by the outer call's own exit, which can
never run because it's blocked inside the inner call waiting for it.
Runs in a background daemon thread with a bounded join so a regression
fails this test instead of hanging the whole suite.
"""
result: dict = {}
def run():
try:
with force_offline_if_cached(False, "outer-uncached"), force_offline_if_cached(True, "inner-cached"):
pass
except Exception as exc:
result["exc"] = exc
else:
result["exc"] = None
t = threading.Thread(target=run, daemon=True)
t.start()
t.join(timeout=3)
assert not t.is_alive(), (
"nesting the opposite mode on the same thread hung instead of raising -- "
"this is the deadlock the per-thread mode-stack guard exists to prevent"
)
assert isinstance(result.get("exc"), RuntimeError), result.get("exc")
if __name__ == "__main__":
pytest.main([__file__, "-v"])
+132 -97
View File
@@ -40,6 +40,24 @@ _saved_env: Optional[str] = None
_saved_hf_const: Optional[bool] = None
_saved_transformers_const: Optional[bool] = None
# Per-thread stack of the modes (`is_cached` values) a thread currently holds
# open. Same-mode nesting on one thread is fine (test_nested_contexts_respect_
# refcount relies on it) and re-enters `_offline_cv`'s RLock without blocking.
# Opposite-mode nesting on the *same* thread would deadlock instead: the inner
# call's `while ...: _offline_cv.wait()` would wait on a condition only some
# *other* thread's exit can clear, but `Condition.wait()` on an RLock releases
# only one recursion level, so the outer call's level stays held and no other
# thread can ever acquire it to notify. Detect and refuse instead of hanging.
_thread_modes = threading.local()
def _mode_stack() -> list:
try:
return _thread_modes.stack
except AttributeError:
_thread_modes.stack = []
return _thread_modes.stack
@contextmanager
def force_offline_if_cached(is_cached: bool, model_label: str = ""):
@@ -57,6 +75,10 @@ def force_offline_if_cached(is_cached: bool, model_label: str = ""):
(or be blamed for breaking) a concurrent cached load's forced-offline
state.
Nesting calls with the *same* ``is_cached`` value on one thread is
supported. Nesting the opposite value on the same thread raises
``RuntimeError`` instead of deadlocking (see ``_thread_modes`` above).
Args:
is_cached: Whether the model weights are already on disk.
model_label: Human-readable name used in log messages.
@@ -64,111 +86,124 @@ def force_offline_if_cached(is_cached: bool, model_label: str = ""):
global _offline_refcount, _uncached_active
global _saved_env, _saved_hf_const, _saved_transformers_const
if not is_cached:
stack = _mode_stack()
if stack and stack[-1] != is_cached:
raise RuntimeError(
f"force_offline_if_cached({is_cached!r}, {model_label!r}) was called "
f"while this thread already holds a force_offline_if_cached({stack[-1]!r}, ...) "
"context open. Nesting the opposite mode on the same thread would "
"deadlock rather than raise; nest same-mode calls only, or run the "
"other mode on a different thread."
)
stack.append(is_cached)
try:
if not is_cached:
with _offline_cv:
while _offline_refcount > 0:
_offline_cv.wait()
_uncached_active += 1
try:
yield
finally:
with _offline_cv:
_uncached_active -= 1
_offline_cv.notify_all()
return
with _offline_cv:
while _offline_refcount > 0:
while _uncached_active > 0:
_offline_cv.wait()
_uncached_active += 1
if _offline_refcount == 0:
# Snapshot prior state, apply new state, roll back on *any*
# failure. Catching only ImportError here would let a partially
# broken install (RuntimeError, AttributeError from a half-init
# module, etc.) leave the cached HF constants mutated without
# bumping the refcount — a persistent offline leak that outlives
# the process and is miserable to debug.
prev_env = os.environ.get("HF_HUB_OFFLINE")
prev_hf: Optional[bool] = None
prev_tf: Optional[bool] = None
try:
try:
import huggingface_hub.constants as hf_const
prev_hf = hf_const.HF_HUB_OFFLINE
hf_const.HF_HUB_OFFLINE = True
except ImportError:
prev_hf = None
try:
import transformers.utils.hub as tf_hub
prev_tf = tf_hub._is_offline_mode
tf_hub._is_offline_mode = True
except ImportError:
prev_tf = None
os.environ["HF_HUB_OFFLINE"] = "1"
except BaseException:
# Roll back whatever we already changed, then re-raise so
# the caller sees the real failure.
if prev_hf is not None:
try:
import huggingface_hub.constants as hf_const
hf_const.HF_HUB_OFFLINE = prev_hf
except ImportError:
pass
if prev_tf is not None:
try:
import transformers.utils.hub as tf_hub
tf_hub._is_offline_mode = prev_tf
except ImportError:
pass
if prev_env is not None:
os.environ["HF_HUB_OFFLINE"] = prev_env
else:
os.environ.pop("HF_HUB_OFFLINE", None)
raise
_saved_env = prev_env
_saved_hf_const = prev_hf
_saved_transformers_const = prev_tf
logger.info(
"[offline-guard] %s is cached — forcing offline mode",
model_label or "model",
)
_offline_refcount += 1
try:
yield
finally:
with _offline_cv:
_uncached_active -= 1
_offline_refcount -= 1
if _offline_refcount == 0:
if _saved_env is not None:
os.environ["HF_HUB_OFFLINE"] = _saved_env
else:
os.environ.pop("HF_HUB_OFFLINE", None)
if _saved_hf_const is not None:
try:
import huggingface_hub.constants as hf_const
hf_const.HF_HUB_OFFLINE = _saved_hf_const
except ImportError:
pass
if _saved_transformers_const is not None:
try:
import transformers.utils.hub as tf_hub
tf_hub._is_offline_mode = _saved_transformers_const
except ImportError:
pass
_saved_env = None
_saved_hf_const = None
_saved_transformers_const = None
_offline_cv.notify_all()
return
with _offline_cv:
while _uncached_active > 0:
_offline_cv.wait()
if _offline_refcount == 0:
# Snapshot prior state, apply new state, roll back on *any*
# failure. Catching only ImportError here would let a partially
# broken install (RuntimeError, AttributeError from a half-init
# module, etc.) leave the cached HF constants mutated without
# bumping the refcount — a persistent offline leak that outlives
# the process and is miserable to debug.
prev_env = os.environ.get("HF_HUB_OFFLINE")
prev_hf: Optional[bool] = None
prev_tf: Optional[bool] = None
try:
try:
import huggingface_hub.constants as hf_const
prev_hf = hf_const.HF_HUB_OFFLINE
hf_const.HF_HUB_OFFLINE = True
except ImportError:
prev_hf = None
try:
import transformers.utils.hub as tf_hub
prev_tf = tf_hub._is_offline_mode
tf_hub._is_offline_mode = True
except ImportError:
prev_tf = None
os.environ["HF_HUB_OFFLINE"] = "1"
except BaseException:
# Roll back whatever we already changed, then re-raise so
# the caller sees the real failure.
if prev_hf is not None:
try:
import huggingface_hub.constants as hf_const
hf_const.HF_HUB_OFFLINE = prev_hf
except ImportError:
pass
if prev_tf is not None:
try:
import transformers.utils.hub as tf_hub
tf_hub._is_offline_mode = prev_tf
except ImportError:
pass
if prev_env is not None:
os.environ["HF_HUB_OFFLINE"] = prev_env
else:
os.environ.pop("HF_HUB_OFFLINE", None)
raise
_saved_env = prev_env
_saved_hf_const = prev_hf
_saved_transformers_const = prev_tf
logger.info(
"[offline-guard] %s is cached — forcing offline mode",
model_label or "model",
)
_offline_refcount += 1
try:
yield
finally:
with _offline_cv:
_offline_refcount -= 1
if _offline_refcount == 0:
if _saved_env is not None:
os.environ["HF_HUB_OFFLINE"] = _saved_env
else:
os.environ.pop("HF_HUB_OFFLINE", None)
if _saved_hf_const is not None:
try:
import huggingface_hub.constants as hf_const
hf_const.HF_HUB_OFFLINE = _saved_hf_const
except ImportError:
pass
if _saved_transformers_const is not None:
try:
import transformers.utils.hub as tf_hub
tf_hub._is_offline_mode = _saved_transformers_const
except ImportError:
pass
_saved_env = None
_saved_hf_const = None
_saved_transformers_const = None
_offline_cv.notify_all()
stack.pop()
_mistral_regex_patched = False