"""Tests for VRAM yield classification and the reclaim-on-OOM path. These cover the two failure modes that motivated the arbitration rework, both of which were observed on real hardware before being encoded here: * A model mid-generation cannot unload. Persisted counters showed 19 "timeouts" in 20 yields; telemetry for that window showed the GPU pinned at 96-97% with 14.92 GB held. That is a busy model, not a fault, and must not be retried in a tight loop. * A model that will not fit fails differently depending on configuration. With n_gpu_layers pinned to 99 (this box) Ollama returns a hard CUDA OOM rather than spilling layers to the CPU: "llama-server process has terminated: exit status 1: cudaMalloc failed: out of memory ... unable to allocate CUDA0 buffer" """ import asyncio import pytest import vram_arbitrator as v GB = 1024 ** 3 def _snap(ollama_gb, util, free_gb=1.0): return {"ollama_bytes": int(ollama_gb * GB), "comfyui_bytes": 0, "other_bytes": 0, "free_bytes": int(free_gb * GB), "gpu_util_pct": util} class TestOomDetection: """The retry path keys off Ollama's error text, so the matcher must be exact.""" def test_matches_the_real_observed_ollama_oom(self): real = ("llama-server process has terminated: exit status 1: cudaMalloc failed: " "out of memory\nalloc_tensor_range: failed to allocate CUDA0 buffer of " "size 13028925440\nerror loading model: unable to allocate CUDA0 buffer") assert v.looks_like_vram_oom(real) @pytest.mark.parametrize("text", [ "cudaMalloc failed: out of memory", "unable to allocate CUDA0 buffer", "failed to allocate buffer", "CUDA error: something", ]) def test_matches_each_signature(self, text): assert v.looks_like_vram_oom(text) @pytest.mark.parametrize("text", [ "model 'foo' not found", "invalid parameter", "", None, "context length exceeded", ]) def test_does_not_match_unrelated_failures(self, text): # A false positive here would purge ComfyUI over a typo in a model name. assert not v.looks_like_vram_oom(text) class TestYieldOutcomeClassification: """_await_vram_release must separate released / busy / stuck.""" def _run(self, snaps, timeout_s=2.0, baseline_gb=14.9): seq = list(snaps) def fake(): return seq.pop(0) if len(seq) > 1 else seq[0] original = v.get_process_vram_bytes v.get_process_vram_bytes = fake try: return asyncio.run( v._await_vram_release(int(baseline_gb * GB), timeout_s=timeout_s)) finally: v.get_process_vram_bytes = original def test_released_when_vram_drains(self): res = self._run([_snap(14.9, 30), _snap(0.0, 5, free_gb=15.4)]) assert res["outcome"] == "released" assert res["confirmed"] is True def test_busy_when_vram_held_and_gpu_pinned(self): # The observed pathology: 14.92 GB held at 96% utilisation. res = self._run([_snap(14.92, 96)]) assert res["outcome"] == "busy" assert res["confirmed"] is False assert "mid-generation" in res["error"] assert res["peak_util_pct"] >= v.BUSY_UTIL_PCT def test_stuck_when_vram_held_and_gpu_idle(self): # VRAM held with nothing running is the genuine fault case. res = self._run([_snap(14.9, 2)], timeout_s=0.3) assert res["outcome"] == "stuck" assert "idle" in res["error"] def test_busy_is_decided_only_after_the_probe_window(self): # Deciding instantly would misread the normal 40-110ms release as busy. assert v.BUSY_PROBE_S > 0 assert v.YIELD_CONFIRM_TIMEOUT_S > v.BUSY_PROBE_S def test_default_wait_is_short(self): # It was 10s, which blocked the arbitrator for the length of an inference while # ComfyUI -- which is not gated on our return value -- waited anyway. assert v.YIELD_CONFIRM_TIMEOUT_S <= 3.0 assert v.YIELD_CONFIRM_TIMEOUT_BLOCKING_S >= 10.0 class TestBusyBackoff: """A busy model must not be re-asked every second.""" def test_backoff_schedule_is_monotonic_and_bounded(self): sched = v.AutoArbitrator.BUSY_BACKOFF_S assert list(sched) == sorted(sched) assert sched[0] >= 1.0 def test_streak_walks_up_the_schedule_and_clamps(self): arb = v.AutoArbitrator() sched = arb.BUSY_BACKOFF_S for streak in range(len(sched) + 3): delay = sched[min(streak, len(sched) - 1)] assert delay == sched[min(streak, len(sched) - 1)] assert sched[min(99, len(sched) - 1)] == sched[-1] def test_release_clears_backoff_state(self): arb = v.AutoArbitrator() arb._yield_backoff_until["m"] = 1e18 arb._yield_busy_streak["m"] = 3 arb.note_deferred_release(1234.0) assert arb._yield_backoff_until == {} assert arb._yield_busy_streak == {} assert arb.stats["deferred_releases"] == 1 def test_counters_distinguish_busy_from_stalled(self): # The old single yield_timeouts counter reported a healthy cron job as a 95% # failure rate. arb = v.AutoArbitrator() assert "yield_deferred_busy" in arb.stats assert "yield_stalled" in arb.stats assert "yield_timeouts" not in arb.stats