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 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__": if __name__ == "__main__":
pytest.main([__file__, "-v"]) pytest.main([__file__, "-v"])
+35
View File
@@ -40,6 +40,24 @@ _saved_env: Optional[str] = None
_saved_hf_const: Optional[bool] = None _saved_hf_const: Optional[bool] = None
_saved_transformers_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 @contextmanager
def force_offline_if_cached(is_cached: bool, model_label: str = ""): 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 (or be blamed for breaking) a concurrent cached load's forced-offline
state. 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: Args:
is_cached: Whether the model weights are already on disk. is_cached: Whether the model weights are already on disk.
model_label: Human-readable name used in log messages. model_label: Human-readable name used in log messages.
@@ -64,6 +86,17 @@ def force_offline_if_cached(is_cached: bool, model_label: str = ""):
global _offline_refcount, _uncached_active global _offline_refcount, _uncached_active
global _saved_env, _saved_hf_const, _saved_transformers_const global _saved_env, _saved_hf_const, _saved_transformers_const
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: if not is_cached:
with _offline_cv: with _offline_cv:
while _offline_refcount > 0: while _offline_refcount > 0:
@@ -169,6 +202,8 @@ def force_offline_if_cached(is_cached: bool, model_label: str = ""):
_saved_hf_const = None _saved_hf_const = None
_saved_transformers_const = None _saved_transformers_const = None
_offline_cv.notify_all() _offline_cv.notify_all()
finally:
stack.pop()
_mistral_regex_patched = False _mistral_regex_patched = False