diff --git a/backend/tests/test_offline_guard.py b/backend/tests/test_offline_guard.py index 1d80ae96..9dc34ab5 100644 --- a/backend/tests/test_offline_guard.py +++ b/backend/tests/test_offline_guard.py @@ -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"]) diff --git a/backend/utils/hf_offline_patch.py b/backend/utils/hf_offline_patch.py index 5b606138..2377f41b 100644 --- a/backend/utils/hf_offline_patch.py +++ b/backend/utils/hf_offline_patch.py @@ -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