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
+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