From e59bb040ee2da2ae09d463bd5e076c4104cb9e6f Mon Sep 17 00:00:00 2001 From: Palash Debnath <4178343+debpalash@users.noreply.github.com> Date: Thu, 1 Oct 2026 18:43:56 +0530 Subject: [PATCH] fix(model): tie abandoned-load verdict to the timed-out worker and clear it when the loader exits (#2394) Review findings on #2467: a late-failing abandoned load left _load_abandoned set; a queued caller's timeout could brand another caller's live load as stuck; the warm-path test had no assertions. Co-Authored-By: Claude Sonnet 5.5 --- backend/services/model_manager.py | 83 ++++++++++++++---- tests/test_cold_load_exclusion_2394.py | 117 ++++++++++++++++++++++--- 2 files changed, 171 insertions(+), 29 deletions(-) diff --git a/backend/services/model_manager.py b/backend/services/model_manager.py index 3734dee85..3076378dd 100644 --- a/backend/services/model_manager.py +++ b/backend/services/model_manager.py @@ -2368,17 +2368,36 @@ class ModelLoadAbandoned(RuntimeError): """ -#: Set by a cold load for as long as it holds ``_model_load_thread_lock`` — -#: i.e. while it is inside the native loader and cannot be cancelled. -#: #2394: this is what lets a later caller tell "another load is still working, -#: wait your turn" (fine, and it will return the model) apart from "a load was -#: already abandoned and may never finish" (a restart is the only way out). -#: An Event, not a flag: the wait must be able to distinguish the two without -#: polling. -_load_in_progress = threading.Event() +class _LoadTicket: + """One caller's claim on a cold load, so a timeout blames the RIGHT load. + + #2394: ``_load_model_with_timeout()`` must mark a load abandoned only when + *its own* worker is the one stuck in the native loader. A process-global + "a load is in progress" flag cannot tell that apart from "my worker is + still queued behind somebody else's load" — it would brand a healthy, + progressing load as abandoned and make every later cold load fail fast. + + ``loading`` is set by the worker once it holds the load lock and is about to + enter the native loader. ``gave_up`` is set by the caller when its deadline + passes. Each side sets its own event BEFORE reading the other's, so either + the caller sees ``loading`` (and abandons the running load) or the worker + sees ``gave_up`` (and declines to start one nobody awaits) — never neither. + """ + + __slots__ = ("loading", "gave_up", "done") + + def __init__(self) -> None: + self.loading = threading.Event() + self.gave_up = threading.Event() + #: Set by the worker as it leaves the critical section, so a caller + #: that times out a hair after the load finished does not brand a + #: finished load as stuck. + self.done = threading.Event() + #: Set when a load is GIVEN UP ON at its deadline while still running (#2394), -#: cleared by that load's own ``finally`` if it ever finishes. Read before +#: cleared by that load's own ``finally`` when it exits — success or failure — +#: because from then on the lock is free and nothing is stuck. Read before #: waiting on the load lock: a caller that arrives after an abandonment must #: fail immediately, because the lock it would queue behind is held by a #: loader nobody is waiting on any more. Without this, a retry re-waits the @@ -2882,7 +2901,9 @@ def _reset_gpu_pool() -> None: _gpu_pool_singleton.reset() -def _load_model_exclusive(timeout: float | None = None): +def _load_model_exclusive( + timeout: float | None = None, ticket: _LoadTicket | None = None +): """The single cold-load leaf: reclaim, load, publish — under ONE lock. #2394. Both cold routes reach the native `from_pretrained` through here: @@ -2947,21 +2968,35 @@ def _load_model_exclusive(timeout: float | None = None): "Logs → Restart backend) if it does not finish on its own, then try " "again." ) - # Reaching the lock clears the "stuck" verdict: whoever held it has - # finished, so the backend is healthy again and needs no restart. A load - # that merely looked stuck from outside is not permanently disarmed. + # Reaching the lock means whoever held it has finished, so any "stuck" + # verdict is stale and the backend needs no restart. _load_abandoned.clear() - _load_in_progress.set() try: + if ticket is not None: + ticket.loading.set() + if ticket.gave_up.is_set(): + # Our caller timed out while we were queued; it already told the + # user and moved on. Starting a native load nobody awaits would + # only pin this lock for the whole load. + raise ModelLoadAbandoned( + "This model load request had already timed out; not " + "starting a load nobody is waiting for." + ) if model is not None: - # A load that finished while we waited. It also clears - # ``_load_in_progress`` on its way out. + # A load that finished while we waited. return model _make_room_before_tts_load() model = _load_model_sync() return model finally: - _load_in_progress.clear() + # Whether this load succeeded, failed or was an abandoned one finishing + # late, it is no longer inside the native loader: the lock is about to + # be free, so nothing is stuck any more. Without this a failed + # abandoned load would leave every later cold load refused until the + # backend restarts, despite a free lock. + if ticket is not None: + ticket.done.set() + _load_abandoned.clear() _model_load_thread_lock.release() @@ -2990,10 +3025,11 @@ async def _load_model_with_timeout(): """ loop = asyncio.get_running_loop() timeout = _model_load_timeout() + ticket = _LoadTicket() try: return await asyncio.wait_for( loop.run_in_executor( - _get_gpu_pool(), _load_model_exclusive, timeout + _get_gpu_pool(), _load_model_exclusive, timeout, ticket ), timeout=timeout, ) @@ -3008,8 +3044,17 @@ async def _load_model_with_timeout(): # rather than re-waiting the full budget behind this loader — the # inline route passes no deadline of its own, so without the flag it # would sit on the lock indefinitely (#2394). - if _load_in_progress.is_set(): + ticket.gave_up.set() + stuck = False + if ticket.loading.is_set(): _load_abandoned.set() + # Re-check after publishing: if the worker finished in between, its + # own clear may already have run and ours would be stale. + if ticket.done.is_set(): + _load_abandoned.clear() + else: + stuck = True + if stuck: logger.error( "Model load exceeded %ss and is still running in a worker that " "cannot be interrupted; the load lock stays held until it " diff --git a/tests/test_cold_load_exclusion_2394.py b/tests/test_cold_load_exclusion_2394.py index 2839c88e1..b6efe1120 100644 --- a/tests/test_cold_load_exclusion_2394.py +++ b/tests/test_cold_load_exclusion_2394.py @@ -255,6 +255,19 @@ def test_a_lone_cold_load_still_loads_and_publishes(mm, monkeypatch): def test_a_warm_model_never_takes_the_load_lock(mm, monkeypatch): """The guard must not serialise the hot path behind a load lock.""" + probe = _ConcurrentLoadProbe() + resident = object() + monkeypatch.setattr(mm, "model", resident, raising=False) + monkeypatch.setattr(mm, "_load_model_sync", probe, raising=False) + + async def _no_heal(): + return None + + monkeypatch.setattr(mm, "_heal_tts_placement", _no_heal, raising=False) + monkeypatch.setattr(mm, "make_room_before_generate", lambda: None, raising=False) + + assert asyncio.run(mm.get_model()) is resident + assert probe.calls == 0, "a resident model must not re-enter the cold load" def test_a_retry_after_a_timed_out_load_must_not_block_forever(mm, monkeypatch): @@ -345,17 +358,101 @@ def _run_capture(mm): except BaseException as exc: # noqa: BLE001 — the failure IS the subject return exc - probe = _ConcurrentLoadProbe() - resident = object() - monkeypatch.setattr(mm, "model", resident, raising=False) - monkeypatch.setattr(mm, "_load_model_sync", probe, raising=False) - async def _no_heal(): - return None +def _wedge_setup(mm, monkeypatch, loader): + from concurrent.futures import ThreadPoolExecutor - monkeypatch.setattr(mm, "_heal_tts_placement", _no_heal, raising=False) - monkeypatch.setattr(mm, "make_room_before_generate", lambda: None, raising=False) + monkeypatch.setattr(mm, "model", None, raising=False) + monkeypatch.setattr(mm, "_model_lock", asyncio.Lock(), raising=False) + monkeypatch.setattr(mm, "_model_load_thread_lock", threading.RLock(), raising=False) + monkeypatch.setattr(mm, "_load_model_sync", loader, raising=False) + monkeypatch.setattr(mm, "_make_room_before_tts_load", lambda: None, raising=False) + monkeypatch.setattr(mm, "_model_load_timeout", lambda: 0.3, raising=False) + mm._load_abandoned.clear() + pool = ThreadPoolExecutor(max_workers=2, thread_name_prefix="test-pool") + monkeypatch.setattr(mm, "_get_gpu_pool", lambda: pool, raising=False) + return pool - assert asyncio.run(mm.get_model()) is resident - assert probe.calls == 0, "a resident model must not re-enter the cold load" +def test_an_abandoned_load_that_later_fails_does_not_strand_the_backend(mm, monkeypatch): + """Review (Greptile P1): the abandoned verdict must die with its loader. + + A timed-out load that eventually FAILS releases the lock but used to leave + ``_load_abandoned`` set, so every later cold load was refused until restart + although nothing was stuck any more. + """ + wedged, release = threading.Event(), threading.Event() + calls = [] + + def _loader(): + calls.append(1) + if len(calls) == 1: + wedged.set() + release.wait(30) + raise RuntimeError("late native load failure") + return object() + + pool = _wedge_setup(mm, monkeypatch, _loader) + try: + first = _run_capture(mm) + assert isinstance(first, mm.ModelLoadAbandoned), first + assert wedged.is_set() + assert mm._load_abandoned.is_set() + + release.set() # the wedged loader now dies on its own + deadline = time.monotonic() + 15 + while mm._load_abandoned.is_set() and time.monotonic() < deadline: + time.sleep(0.01) + assert not mm._load_abandoned.is_set(), ( + "the abandoned flag outlived its loader; the backend stays refused" + ) + monkeypatch.setattr(mm, "_get_gpu_pool", lambda: pool, raising=False) + retry = _run_capture(mm) + assert not isinstance(retry, BaseException), retry + finally: + release.set() + monkeypatch.setattr(mm, "model", None, raising=False) + pool.shutdown(wait=False) + + +def test_a_caller_queued_behind_another_load_does_not_brand_it_abandoned(mm, monkeypatch): + """Review (Greptile P1): blame the worker that timed out, not whichever + load happens to be running. + + Caller B times out while still queued behind caller A's healthy, in-flight + load. A's load must not be marked abandoned. + """ + in_native, release = threading.Event(), threading.Event() + published = object() + + def _loader(): + in_native.set() + release.wait(30) + return published + + pool = _wedge_setup(mm, monkeypatch, _loader) + monkeypatch.setattr(mm, "_model_load_timeout", lambda: 5.0, raising=False) + try: + # A: an inline-route load that holds the lock with no deadline. + a_result: dict[str, object] = {} + a = threading.Thread( + target=lambda: a_result.setdefault("r", mm._load_model_exclusive()), + daemon=True, + ) + a.start() + assert in_native.wait(10) + + # B: a pool route whose deadline (shortened) passes while queued. + monkeypatch.setattr(mm, "_model_load_timeout", lambda: 0.3, raising=False) + b = _run_capture(mm) + assert isinstance(b, RuntimeError), b + assert not mm._load_abandoned.is_set(), ( + "a queued caller's timeout branded another caller's live load as stuck" + ) + release.set() + a.join(timeout=15) + assert a_result.get("r") is published + finally: + release.set() + monkeypatch.setattr(mm, "model", None, raising=False) + pool.shutdown(wait=False)