|
| 1 | +"""Unit tests for the boot-time warm-up logic in handler.py. |
| 2 | +
|
| 3 | +These need no GPU, no network and no running Ollama, which is the point: the |
| 4 | +warm-up only ever runs on real hardware at cold start, so getting its request |
| 5 | +shape, bounds or failure handling wrong is otherwise only discovered by a timed |
| 6 | +out first request in production. |
| 7 | +
|
| 8 | + python3 -m pytest test_warmup.py -v |
| 9 | +""" |
| 10 | + |
| 11 | +import json |
| 12 | +import sys |
| 13 | +import types |
| 14 | + |
| 15 | +import pytest |
| 16 | + |
| 17 | +# handler.py imports the runpod SDK at module scope; stub it when it isn't |
| 18 | +# installed so the pure functions stay testable from a bare checkout. |
| 19 | +if "runpod" not in sys.modules: |
| 20 | + try: |
| 21 | + import runpod # noqa: F401 |
| 22 | + except ImportError: |
| 23 | + sys.modules["runpod"] = types.ModuleType("runpod") |
| 24 | + |
| 25 | +import requests |
| 26 | + |
| 27 | +import handler |
| 28 | + |
| 29 | + |
| 30 | +class FakeResponse: |
| 31 | + def __init__(self, status_code=200, payload=None, text=None, path="/api/generate"): |
| 32 | + self.status_code = status_code |
| 33 | + self.ok = status_code < 400 |
| 34 | + self._payload = payload |
| 35 | + self.text = text if text is not None else json.dumps(payload or {}) |
| 36 | + self.request = types.SimpleNamespace(path_url=path) |
| 37 | + |
| 38 | + def json(self): |
| 39 | + if self._payload is None: |
| 40 | + raise ValueError("no json") |
| 41 | + return self._payload |
| 42 | + |
| 43 | + |
| 44 | +class FakeSession: |
| 45 | + """Records every call so tests can assert on the exact request shape.""" |
| 46 | + |
| 47 | + def __init__(self, post_response, ps_response=None): |
| 48 | + self.post_response = post_response |
| 49 | + self.ps_response = ps_response or FakeResponse(payload={"models": []}, path="/api/ps") |
| 50 | + self.posts = [] |
| 51 | + self.gets = [] |
| 52 | + |
| 53 | + def post(self, url, **kwargs): |
| 54 | + self.posts.append((url, kwargs)) |
| 55 | + return self.post_response |
| 56 | + |
| 57 | + def get(self, url, **kwargs): |
| 58 | + self.gets.append((url, kwargs)) |
| 59 | + return self.ps_response |
| 60 | + |
| 61 | + |
| 62 | +def loaded_ps(model, size=1000, size_vram=None): |
| 63 | + return FakeResponse( |
| 64 | + payload={ |
| 65 | + "models": [ |
| 66 | + {"name": model, "size": size, "size_vram": size if size_vram is None else size_vram} |
| 67 | + ] |
| 68 | + }, |
| 69 | + path="/api/ps", |
| 70 | + ) |
| 71 | + |
| 72 | + |
| 73 | +# --- _duration_seconds: Go-style durations, as Ollama's env vars use ------------ |
| 74 | + |
| 75 | + |
| 76 | +@pytest.mark.parametrize( |
| 77 | + "value, expected", |
| 78 | + [ |
| 79 | + ("60m", 3600.0), |
| 80 | + ("1h30m", 5400.0), |
| 81 | + ("90s", 90.0), |
| 82 | + ("500ms", 0.5), |
| 83 | + ("1.5h", 5400.0), |
| 84 | + ("-5m", -300.0), |
| 85 | + ], |
| 86 | +) |
| 87 | +def test_duration_seconds_parses(value, expected): |
| 88 | + assert handler._duration_seconds(value, 123.0) == pytest.approx(expected) |
| 89 | + |
| 90 | + |
| 91 | +@pytest.mark.parametrize("value", ["", None, "garbage", "10", "5x", "m5", "1h 30m"]) |
| 92 | +def test_duration_seconds_falls_back(value): |
| 93 | + assert handler._duration_seconds(value, 123.0) == 123.0 |
| 94 | + |
| 95 | + |
| 96 | +# --- _warmup_timeout_seconds: always bounded, tracks OLLAMA_LOAD_TIMEOUT -------- |
| 97 | + |
| 98 | + |
| 99 | +def test_warmup_timeout_tracks_load_timeout(monkeypatch): |
| 100 | + monkeypatch.setenv("OLLAMA_LOAD_TIMEOUT", "10m") |
| 101 | + assert handler._warmup_timeout_seconds() == 600 + handler._WARMUP_MARGIN_S |
| 102 | + |
| 103 | + |
| 104 | +@pytest.mark.parametrize("value", [None, "0s", "-1s", "not-a-duration"]) |
| 105 | +def test_warmup_timeout_never_unbounded(monkeypatch, value): |
| 106 | + if value is None: |
| 107 | + monkeypatch.delenv("OLLAMA_LOAD_TIMEOUT", raising=False) |
| 108 | + else: |
| 109 | + monkeypatch.setenv("OLLAMA_LOAD_TIMEOUT", value) |
| 110 | + assert ( |
| 111 | + handler._warmup_timeout_seconds() |
| 112 | + == handler._WARMUP_FALLBACK_S + handler._WARMUP_MARGIN_S |
| 113 | + ) |
| 114 | + |
| 115 | + |
| 116 | +# --- warm_model: the load request itself ---------------------------------------- |
| 117 | + |
| 118 | + |
| 119 | +def test_warm_model_sends_empty_prompt_load(monkeypatch): |
| 120 | + monkeypatch.setenv("OLLAMA_LOAD_TIMEOUT", "10m") |
| 121 | + session = FakeSession( |
| 122 | + FakeResponse(payload={"done": True}), ps_response=loaded_ps("llama3.2:3b") |
| 123 | + ) |
| 124 | + monkeypatch.setattr(handler, "session", session) |
| 125 | + |
| 126 | + handler.warm_model("llama3.2:3b") |
| 127 | + |
| 128 | + url, kwargs = session.posts[0] |
| 129 | + assert url.endswith("/api/generate") |
| 130 | + # An empty prompt is Ollama's "load into memory, generate nothing" request. |
| 131 | + assert kwargs["json"] == {"model": "llama3.2:3b", "prompt": "", "stream": False} |
| 132 | + assert kwargs["timeout"] == 600 + handler._WARMUP_MARGIN_S |
| 133 | + |
| 134 | + |
| 135 | +def test_warm_model_raises_on_server_error(monkeypatch): |
| 136 | + session = FakeSession( |
| 137 | + FakeResponse( |
| 138 | + status_code=500, |
| 139 | + payload={"error": "timed out waiting for llama runner to start"}, |
| 140 | + ) |
| 141 | + ) |
| 142 | + monkeypatch.setattr(handler, "session", session) |
| 143 | + |
| 144 | + with pytest.raises(ValueError, match="timed out waiting for llama runner"): |
| 145 | + handler.warm_model("llama3.2:3b") |
| 146 | + |
| 147 | + |
| 148 | +def test_warm_model_warns_on_cpu_offload(monkeypatch, capsys): |
| 149 | + session = FakeSession( |
| 150 | + FakeResponse(payload={"done": True}), |
| 151 | + ps_response=loaded_ps("big-model:latest", size=1000, size_vram=400), |
| 152 | + ) |
| 153 | + monkeypatch.setattr(handler, "session", session) |
| 154 | + |
| 155 | + handler.warm_model("big-model") |
| 156 | + |
| 157 | + out = capsys.readouterr().out |
| 158 | + assert "offloaded to CPU" in out |
| 159 | + |
| 160 | + |
| 161 | +def test_warm_model_reports_full_residency(monkeypatch, capsys): |
| 162 | + session = FakeSession( |
| 163 | + FakeResponse(payload={"done": True}), ps_response=loaded_ps("llama3.2:3b") |
| 164 | + ) |
| 165 | + monkeypatch.setattr(handler, "session", session) |
| 166 | + |
| 167 | + handler.warm_model("llama3.2:3b") |
| 168 | + |
| 169 | + assert "resident in GPU memory" in capsys.readouterr().out |
| 170 | + |
| 171 | + |
| 172 | +def test_warm_model_survives_ps_failure(monkeypatch, capsys): |
| 173 | + """A broken /api/ps must not turn a successful load into a failure.""" |
| 174 | + session = FakeSession(FakeResponse(payload={"done": True})) |
| 175 | + |
| 176 | + def broken_get(url, **kwargs): |
| 177 | + raise requests.ConnectionError("ps is down") |
| 178 | + |
| 179 | + session.get = broken_get |
| 180 | + monkeypatch.setattr(handler, "session", session) |
| 181 | + |
| 182 | + handler.warm_model("llama3.2:3b") # must not raise |
| 183 | + |
| 184 | + assert "could not read /api/ps" in capsys.readouterr().out |
| 185 | + |
| 186 | + |
| 187 | +# --- warm_default_model: non-fatal, but says what failed ------------------------- |
| 188 | + |
| 189 | + |
| 190 | +def test_warm_default_model_success(monkeypatch): |
| 191 | + monkeypatch.setattr(handler, "resolve_default_model", lambda: "llama3.2:3b") |
| 192 | + warmed = [] |
| 193 | + monkeypatch.setattr(handler, "warm_model", warmed.append) |
| 194 | + |
| 195 | + assert handler.warm_default_model() is True |
| 196 | + assert warmed == ["llama3.2:3b"] |
| 197 | + |
| 198 | + |
| 199 | +@pytest.mark.parametrize( |
| 200 | + "err", |
| 201 | + [ |
| 202 | + requests.ConnectionError("connection refused"), |
| 203 | + requests.Timeout("read timed out"), |
| 204 | + ValueError("HTTP 500 from /api/generate: model requires more system memory"), |
| 205 | + ], |
| 206 | +) |
| 207 | +def test_warm_default_model_failure_is_nonfatal_but_loud(monkeypatch, capsys, err): |
| 208 | + monkeypatch.setattr(handler, "resolve_default_model", lambda: "big-model") |
| 209 | + |
| 210 | + def failing_warm(model): |
| 211 | + raise err |
| 212 | + |
| 213 | + monkeypatch.setattr(handler, "warm_model", failing_warm) |
| 214 | + |
| 215 | + assert handler.warm_default_model() is False |
| 216 | + out = capsys.readouterr().out |
| 217 | + assert "could not pre-load 'big-model'" in out |
| 218 | + # The cause must be surfaced, not swallowed behind a bare WARN. |
| 219 | + assert str(err) in out |
0 commit comments