"""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 fake_tokenizer import MergeTokenizer, messages, upstream_prefix 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) # The shape of SemIf's score_shared, compiled INTO the fake module so that, like the real one, it # resolves _state_prefix and direct_messages as module globals at call time and refuses (ValueError) # a prefix that is not a token-prefix of every row. SHARED_SRC = """ def score_shared(model, tokenizer, rows, metadata, max_tokens=4096): prefix = _state_prefix(tokenizer, rows[0]["state"]) for row in rows: ids = tokenizer.encode(tokenizer.apply_chat_template(direct_messages(row), tokenize=False, add_generation_prompt=True, enable_thinking=False), add_special_tokens=False) if not prefix or ids[:len(prefix)] != prefix: raise ValueError("The fixed state prefix does not match every full prompt") CALLS.append(("shared", rows[0]["state"])) return [{"id": row["id"]} for row in rows], {} """ # A SemIf revision that binds the helper when score_shared is defined: the module global can be # replaced all day and scoring never sees it. SHARED_SRC_EARLY_BOUND = SHARED_SRC.replace("max_tokens=4096):", "max_tokens=4096, _state_prefix=_state_prefix):") @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"]), MergeTokenizer(), {"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 core.direct_messages = messages direct = types.ModuleType("semif_phase1.direct") direct.score = score shared = types.ModuleType("semif_phase1.shared") shared._state_prefix, shared.direct_messages, shared.CALLS = upstream_prefix, messages, calls exec(SHARED_SRC, shared.__dict__) 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", "shared"] 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" def test_load_wraps_the_upstream_state_prefix_once_even_across_reloads(fakes): """INV-7: score_shared looks _state_prefix up as a module global, so the wrapper must replace it there. A second load must wrap the ORIGINAL again, not stack wrapper on wrapper.""" import semif_phase1.shared as upstream from semif_serve.engine import TorchEngine for _ in range(2): TorchEngine.load(Settings(api_token=TOKEN)) assert upstream._state_prefix is not upstream_prefix assert upstream._state_prefix.semif_serve_wraps is upstream_prefix def test_load_fails_closed_when_the_pinned_semif_no_longer_has_the_prefix_hook(fakes, monkeypatch): import semif_phase1.shared as upstream from semif_serve.engine import TorchEngine _torch, calls, _ = fakes monkeypatch.delattr(upstream, "_state_prefix") with pytest.raises(RuntimeError, match="_state_prefix"): TorchEngine.load(Settings(api_token=TOKEN)) assert not any(c[0] == "load" for c in calls) def test_load_proves_the_hook_on_a_merge_prone_state_through_the_real_shared_path(fakes): """Bug hunt SKAL (R2/H1, H2): existence is not effect. Startup scores one state that upstream alone refuses, through score_shared itself.""" import semif_phase1.shared as upstream from semif_serve.engine import BOUNDARY_ROW, TorchEngine _torch, calls, _ = fakes with pytest.raises(ValueError): # the probe really is merge-prone here upstream.score_shared(None, MergeTokenizer(), [BOUNDARY_ROW], {}) TorchEngine.load(Settings(api_token=TOKEN)) assert ("shared", BOUNDARY_ROW["state"]) in calls def test_load_fails_closed_when_score_shared_binds_the_prefix_before_the_patch(fakes): import semif_phase1.shared as upstream from semif_serve.engine import TorchEngine exec(SHARED_SRC_EARLY_BOUND, upstream.__dict__) with pytest.raises(RuntimeError, match="INV-7"): TorchEngine.load(Settings(api_token=TOKEN)) def test_load_fails_closed_when_score_shared_resolves_its_globals_elsewhere(fakes): import semif_phase1.shared as upstream from semif_serve.engine import TorchEngine _torch, calls, _ = fakes elsewhere = {"_state_prefix": upstream_prefix, "direct_messages": messages, "CALLS": calls} exec(SHARED_SRC, elsewhere) upstream.score_shared = elsewhere["score_shared"] # re-exported from another module with pytest.raises(RuntimeError, match="INV-7"): TorchEngine.load(Settings(api_token=TOKEN)) assert not any(c[0] == "load" for c in calls) def test_load_fails_closed_when_the_hook_is_not_callable(fakes, monkeypatch): import semif_phase1.shared as upstream from semif_serve.engine import TorchEngine _torch, calls, _ = fakes monkeypatch.setattr(upstream, "_state_prefix", "not a function") with pytest.raises(RuntimeError, match="_state_prefix"): TorchEngine.load(Settings(api_token=TOKEN)) assert not any(c[0] == "load" for c in calls) def test_load_fails_closed_when_the_wrapper_renders_the_wrong_prompt(fakes, monkeypatch): """The SURVIVED row of bug hunt SKAL: a wrapper that renders nothing still yields a valid (degenerate) prefix, so scoring stays correct and silently loses all sharing. Startup requires the wrapper to keep upstream's whole prefix on an ordinary state.""" import semif_phase1.core as core from semif_serve.engine import TorchEngine monkeypatch.setattr(core, "direct_messages", lambda row: []) with pytest.raises(RuntimeError, match="INV-7"): TorchEngine.load(Settings(api_token=TOKEN))