"""TorchEngine.load and the entry point, against fake torch/semif modules (INV-3, INV-4, INV-5). The bug hunt found load() had no test at all: removing the arch check survived every test.""" import os import sys import types import pytest from semif_serve.config import Settings TOKEN = "t" * 40 class Param: def __init__(self, device): self.device = types.SimpleNamespace(type=device) class Model: def __init__(self, device): self._device = device def parameters(self): yield Param(self._device) @pytest.fixture def fakes(monkeypatch): calls = [] cuda = types.SimpleNamespace( is_available=lambda: True, get_device_capability=lambda _i=0: (12, 0), get_arch_list=lambda: ["sm_90", "sm_120"], get_device_properties=lambda _i=0: types.SimpleNamespace(total_memory=96 * 2**30), set_per_process_memory_fraction=lambda f, _i=0: calls.append(("cap", round(f, 4))), memory_reserved=lambda _i=0: 8 * 2**30, OutOfMemoryError=type("OutOfMemoryError", (RuntimeError,), {}), empty_cache=lambda: calls.append(("empty",)), ) torch = types.SimpleNamespace(cuda=cuda, __version__="2.10.0+cu128") state = {"device": "cuda", "warmup_raises": None} def load_causal_model(model, revision, device, dtype): calls.append(("load", model, revision, device, dtype)) return Model(state["device"]), object(), {"source": model} def score(model, tok, row, meta, max_tokens): calls.append(("score", row["id"])) if state["warmup_raises"]: raise state["warmup_raises"] return {"id": row["id"]} core = types.ModuleType("semif_phase1.core") core.load_causal_model = load_causal_model direct = types.ModuleType("semif_phase1.direct") direct.score = score shared = types.ModuleType("semif_phase1.shared") shared.score_shared = lambda *a: ([], {}) pkg = types.ModuleType("semif_phase1") for name, mod in {"torch": torch, "semif_phase1": pkg, "semif_phase1.core": core, "semif_phase1.direct": direct, "semif_phase1.shared": shared}.items(): monkeypatch.setitem(sys.modules, name, mod) return torch, calls, state def test_load_caps_before_the_weights_land_then_warms_up(fakes): from semif_serve.engine import TorchEngine _torch, calls, _ = fakes TorchEngine.load(Settings(api_token=TOKEN, vram_cap_gib=12.0)) assert [c[0] for c in calls] == ["cap", "load", "score"] assert calls[0] == ("cap", round(12 / 96, 4)) assert calls[2] == ("score", "semif-serve-warmup") def test_load_refuses_a_card_torch_has_no_kernels_for(fakes): from semif_serve.engine import TorchEngine torch, calls, _ = fakes torch.cuda.get_arch_list = lambda: ["sm_80", "sm_90"] with pytest.raises(RuntimeError, match="sm_120"): TorchEngine.load(Settings(api_token=TOKEN)) assert not any(c[0] == "load" for c in calls) def test_load_refuses_a_model_that_landed_on_the_wrong_device(fakes): from semif_serve.engine import TorchEngine _torch, _calls, state = fakes state["device"] = "cpu" with pytest.raises(RuntimeError, match="landed on cpu"): TorchEngine.load(Settings(api_token=TOKEN)) def test_load_fails_closed_when_the_warmup_decision_fails(fakes): from semif_serve.engine import TorchEngine from semif_serve.errors import ScoringFailed _torch, _calls, state = fakes state["warmup_raises"] = RuntimeError("Failed to find C compiler") with pytest.raises(ScoringFailed, match="C compiler"): TorchEngine.load(Settings(api_token=TOKEN)) def test_the_entry_point_forces_offline_mode_before_the_engine_loads(fakes, monkeypatch): import semif_serve.engine as engine_mod from semif_serve import main seen = {} monkeypatch.delenv("HF_HUB_OFFLINE", raising=False) monkeypatch.setenv("SEMIF_API_TOKEN", TOKEN) monkeypatch.setattr(engine_mod.TorchEngine, "load", classmethod(lambda cls, s: seen.update(offline=os.environ.get("HF_HUB_OFFLINE")) or object())) main.app_from_env() assert seen["offline"] == "1"