unsloth/studio/backend/tests/test_gguf_load_cache_reuse.py
Daniel Han d7cdc96051
studio/tests: cover the GGUF load ordering behaviourally and make the structlog stub order-independent (#7442)
* studio: fix Backend CI red on main from an ambiguous ordering anchor

test_load_marker_precedes_hub_guard_and_unload fails on main, so every
open PR against the repo inherits the failure.

Root cause. #7239 (a7761e174) reworked the GGUF GPU-pool validation in
_load_model_impl from "if config.is_gguf and effective_gpu_ids is not
None:" to a bare "if config.is_gguf:", placed earlier in the function
than the GGUF load branch. The test anchors on
source.index("if config.is_gguf:"), a first-match search, so it silently
re-anchored onto the GPU-pool statement. #7251 (95f42bcce) then restored
the assertion "= _resolve_inherited_extra_args(" before
"if config.is_gguf:" against a tree where that anchor already pointed at
the wrong statement, and main went red. Checking out 95f42bcce and
running the suite reproduces the same single failure.

The code is correct. _resolve_inherited_extra_args still runs before the
GGUF load branch and before the hub-download guard that consumes
extra_llama_args for require_mmproj, so the guarantee #7251 protects is
intact; only the assertion is wrong.

Fix. Assert that guarantee behaviourally instead of by source offsets.
The new test drives _load_model_impl over a vision GGUF with a stored
--no-mmproj from a previous same-model load and captures the
require_mmproj the hub guard is called with: inherited --no-mmproj gives
False, nothing to inherit gives True, and an explicit request list wins
over the stored one both ways. Moving the resolution call after the
guard makes the inherited case report True and the test fails, so it
detects the reorder the old assertion was meant to catch, without
depending on how many "if config.is_gguf:" statements the endpoint has.

The surviving marker-before-guard-before-unload assertion had the same
ambiguous anchor for its slice start, silently widening the slice past
the GPU-pool block. It now slices from the "if config.is_gguf:" nearest
above the in-flight marker, which pins the load branch.

The structlog test stub gains a get_logger factory so routes/inference.py
is importable when structlog is absent.

34 pass in tests/test_gguf_load_cache_reuse.py (was 32 pass, 1 fail);
350 pass across it plus test_llama_cpp_mmproj_fallback.py and
test_llama_cpp_mtp_detection.py. A full backend run before and after is
identical apart from this test going from fail to pass.

* studio/tests: repair a pre-existing bare structlog stub before importing routes

* studio/tests: tighten the comments on the new load-ordering coverage

* Tighten comments on the load-ordering coverage for PR #7442
2026-07-26 05:01:56 -07:00

949 lines
35 KiB
Python

# SPDX-License-Identifier: AGPL-3.0-only
# Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. See /studio/LICENSE.AGPL-3.0
"""Tests for cached GGUF reuse and load/download exclusion.
No GPU, network, or subprocesses are required.
"""
from __future__ import annotations
import asyncio
import importlib.util
import logging
import sys
import threading
import types as _types
from contextlib import nullcontext
from pathlib import Path
from types import SimpleNamespace
from unittest.mock import patch
import pytest
_BACKEND_DIR = str(Path(__file__).resolve().parent.parent)
if _BACKEND_DIR not in sys.path:
sys.path.insert(0, _BACKEND_DIR)
# Stub optional dependencies before importing the modules under test.
_loggers_stub = _types.ModuleType("loggers")
_loggers_stub.get_logger = lambda name: __import__("logging").getLogger(name)
sys.modules.setdefault("loggers", _loggers_stub)
_structlog_stub = _types.ModuleType("structlog")
# routes/inference.py binds structlog.get_logger at import time, and setdefault
# keeps a bare stub an earlier test left behind: repair it rather than rely on order.
_structlog_stub.get_logger = lambda *_args, **_kwargs: logging.getLogger("structlog_stub")
sys.modules.setdefault("structlog", _structlog_stub)
if not hasattr(sys.modules["structlog"], "get_logger"):
sys.modules["structlog"].get_logger = _structlog_stub.get_logger
try:
import httpx # noqa: F401
except ImportError:
_httpx_stub = _types.ModuleType("httpx")
for _exc_name in (
"ConnectError",
"TimeoutException",
"ReadTimeout",
"ReadError",
"RemoteProtocolError",
"CloseError",
"HTTPError",
"RequestError",
"HTTPStatusError",
):
setattr(_httpx_stub, _exc_name, type(_exc_name, (Exception,), {}))
_httpx_stub.Response = type("Response", (), {})
_httpx_stub.Request = type("Request", (), {})
class _FakeTimeout:
def __init__(self, *a, **kw):
pass
_httpx_stub.Timeout = _FakeTimeout
_httpx_stub.Client = type(
"Client",
(),
{
"__init__": lambda self, **kw: None,
"__enter__": lambda self: self,
"__exit__": lambda self, *a: None,
},
)
sys.modules.setdefault("httpx", _httpx_stub)
from huggingface_hub import constants as hf_constants
from core.inference.llama_cpp import (
LlamaCppBackend,
cached_gguf_for_load,
gguf_load_in_flight,
hf_gguf_load_in_flight,
)
REPO = "unsloth/gemma-test-GGUF"
VARIANT = "UD-Q4_K_XL"
MAIN = f"gemma-test-{VARIANT}.gguf"
def _build_cache(
root: Path,
repo_id: str,
files: dict[str, int],
*,
snapshot_sha: str = "a" * 40,
) -> Path:
"""Create ``$root/models--<repo>/snapshots/<sha>/<rel>`` for each entry."""
repo_dir = root / f"models--{repo_id.replace('/', '--')}"
(repo_dir / "blobs").mkdir(parents = True, exist_ok = True)
snap = repo_dir / "snapshots" / snapshot_sha
snap.mkdir(parents = True, exist_ok = True)
for rel, size in files.items():
full = snap / rel
full.parent.mkdir(parents = True, exist_ok = True)
full.write_bytes(b"\0" * size)
return snap
@pytest.fixture
def hf_cache(tmp_path, monkeypatch):
monkeypatch.setattr(hf_constants, "HF_HUB_CACHE", str(tmp_path))
monkeypatch.setattr(
"utils.hf_cache_settings.get_hf_cache_paths",
lambda: _types.SimpleNamespace(hub_cache = tmp_path),
)
monkeypatch.delenv("HF_HUB_OFFLINE", raising = False)
monkeypatch.delenv("TRANSFORMERS_OFFLINE", raising = False)
return tmp_path
def _fail_download(*_args, **_kwargs):
raise AssertionError("must reuse the cached GGUF instead of downloading")
def _fail_get_paths_info(*_args, **_kwargs):
raise AssertionError("cached reuse must return before the sizing preflight")
def _load_route_module(name: str, relative_path: str):
"""Import a route module under a private name so patches can't leak."""
spec = importlib.util.spec_from_file_location(name, Path(_BACKEND_DIR) / relative_path)
module = importlib.util.module_from_spec(spec)
spec.loader.exec_module(module)
return module
async def _inline_to_thread(func, /, *args, **kwargs):
return func(*args, **kwargs)
async def _no_gguf_gpu_ids(*_args, **_kwargs):
return None
class TestLoadReusesCachedCopy:
def test_download_uses_selected_cache_for_lookup_preflight_and_write(
self, tmp_path, monkeypatch
):
backend = LlamaCppBackend()
selected = tmp_path / "selected" / "hub"
startup = tmp_path / "startup" / "hub"
monkeypatch.setattr(hf_constants, "HF_HUB_CACHE", str(startup))
monkeypatch.setattr(
"utils.hf_cache_settings.get_hf_cache_paths",
lambda: _types.SimpleNamespace(hub_cache = selected),
)
seen = {"lookups": [], "disk": [], "downloads": []}
def cached_lookup(
repo_id,
filename,
*,
cache_dir = None,
**_kwargs,
):
seen["lookups"].append((repo_id, filename, cache_dir))
return None
def disk_usage(path):
seen["disk"].append(str(path))
return _types.SimpleNamespace(free = 1024)
def download(repo_id, filename, _token, **kwargs):
seen["downloads"].append((repo_id, filename, kwargs.get("cache_dir")))
return str(selected / filename)
with (
patch("huggingface_hub.list_repo_files", lambda *_a, **_k: [MAIN]),
patch(
"huggingface_hub.get_paths_info",
lambda _repo, paths, **_kwargs: [
_types.SimpleNamespace(path = path, size = 4) for path in paths
],
),
patch("huggingface_hub.try_to_load_from_cache", cached_lookup),
patch("core.inference.llama_cpp.shutil.disk_usage", disk_usage),
patch(
"core.inference.llama_cpp.hf_hub_download_with_xet_fallback",
download,
),
):
out = backend._download_gguf(hf_repo = REPO, hf_variant = VARIANT)
assert out == str(selected / MAIN)
assert seen == {
"lookups": [(REPO, MAIN, str(selected))],
"disk": [str(selected)],
"downloads": [(REPO, MAIN, str(selected))],
}
def test_online_reuse_after_revision_bump(self, hf_cache):
"""A new repo revision does not replace a complete cached model."""
backend = LlamaCppBackend()
snap = _build_cache(hf_cache, REPO, {MAIN: 4})
with (
patch("huggingface_hub.list_repo_files", lambda *_a, **_k: [MAIN]),
patch("huggingface_hub.get_paths_info", _fail_get_paths_info),
patch("core.inference.llama_cpp.hf_hub_download_with_xet_fallback", _fail_download),
):
out = backend._download_gguf(hf_repo = REPO, hf_variant = VARIANT)
assert out == str(snap / MAIN)
def test_reuse_size_check_uses_cached_snapshot_revision(self, hf_cache):
"""Current-revision size changes do not invalidate an older complete copy."""
backend = LlamaCppBackend()
snap = _build_cache(hf_cache, REPO, {MAIN: 4})
revisions: list[str | None] = []
def fake_get_paths_info(
_repo,
paths,
*,
revision = None,
token = None,
):
revisions.append(revision)
size = 4 if revision == snap.name else 8
return [_types.SimpleNamespace(path = path, size = size) for path in paths]
with (
patch("huggingface_hub.list_repo_files", lambda *_a, **_k: [MAIN]),
patch("huggingface_hub.get_paths_info", fake_get_paths_info),
patch("core.inference.llama_cpp.hf_hub_download_with_xet_fallback", _fail_download),
):
out = backend._download_gguf(hf_repo = REPO, hf_variant = VARIANT)
assert out == str(snap / MAIN)
assert revisions == [snap.name]
def test_reuse_when_cached_revision_vanished_from_hub(self, hf_cache):
"""The Hub answers an unknown revision with an empty result, not an error."""
backend = LlamaCppBackend()
snap = _build_cache(hf_cache, REPO, {MAIN: 4})
with (
patch("huggingface_hub.list_repo_files", lambda *_a, **_k: [MAIN]),
patch("huggingface_hub.get_paths_info", lambda *_a, **_k: []),
patch("core.inference.llama_cpp.hf_hub_download_with_xet_fallback", _fail_download),
):
out = backend._download_gguf(hf_repo = REPO, hf_variant = VARIANT)
assert out == str(snap / MAIN)
def test_truncated_cached_file_is_not_reused(self, hf_cache):
backend = LlamaCppBackend()
_build_cache(hf_cache, REPO, {MAIN: 4})
downloaded: list[str] = []
def fake_get_paths_info(
_repo,
paths,
*,
revision = None,
token = None,
):
return [_types.SimpleNamespace(path = path, size = 8) for path in paths]
def fake_download(
repo_id,
filename,
token = None,
**_kwargs,
):
downloaded.append(filename)
return f"/fake/{repo_id}/{filename}"
with (
patch("huggingface_hub.list_repo_files", lambda *_a, **_k: [MAIN]),
patch("huggingface_hub.get_paths_info", fake_get_paths_info),
patch("huggingface_hub.try_to_load_from_cache", lambda *_a, **_k: None),
patch("core.inference.llama_cpp.hf_hub_download_with_xet_fallback", fake_download),
):
out = backend._download_gguf(hf_repo = REPO, hf_variant = VARIANT)
assert downloaded == [MAIN]
assert out == f"/fake/{REPO}/{MAIN}"
def test_truncated_cached_split_shard_is_not_reused(self, hf_cache):
backend = LlamaCppBackend()
shard1 = f"gemma-test-{VARIANT}-00001-of-00002.gguf"
shard2 = f"gemma-test-{VARIANT}-00002-of-00002.gguf"
_build_cache(hf_cache, REPO, {shard1: 8, shard2: 4})
downloaded: list[str] = []
def fake_get_paths_info(
_repo,
paths,
*,
revision = None,
token = None,
):
return [_types.SimpleNamespace(path = path, size = 8) for path in paths]
def fake_download(
repo_id,
filename,
token = None,
**_kwargs,
):
downloaded.append(filename)
return f"/fake/{repo_id}/{filename}"
with (
patch("huggingface_hub.list_repo_files", lambda *_a, **_k: [shard1, shard2]),
patch("huggingface_hub.get_paths_info", fake_get_paths_info),
patch("huggingface_hub.try_to_load_from_cache", lambda *_a, **_k: None),
patch("core.inference.llama_cpp.hf_hub_download_with_xet_fallback", fake_download),
):
out = backend._download_gguf(hf_repo = REPO, hf_variant = VARIANT)
assert downloaded == [shard1, shard2]
assert out == f"/fake/{REPO}/{shard1}"
def test_online_reuse_when_reupload_renamed_the_file(self, hf_cache):
"""A renamed variant still reuses its cached file."""
backend = LlamaCppBackend()
old_name = f"gemma-test-old-{VARIANT}.gguf"
snap = _build_cache(hf_cache, REPO, {old_name: 4})
with (
patch("huggingface_hub.list_repo_files", lambda *_a, **_k: [MAIN]),
patch("huggingface_hub.get_paths_info", _fail_get_paths_info),
patch("core.inference.llama_cpp.hf_hub_download_with_xet_fallback", _fail_download),
):
out = backend._download_gguf(hf_repo = REPO, hf_variant = VARIANT)
assert out == str(snap / old_name)
def test_downloads_when_nothing_cached(self, hf_cache):
backend = LlamaCppBackend()
downloaded: list[str] = []
def fake_download(
repo_id,
filename,
token = None,
**_kwargs,
):
downloaded.append(filename)
return f"/fake/{repo_id}/{filename}"
def fake_get_paths_info(
_repo_id,
paths,
token = None,
):
return [_types.SimpleNamespace(path = p, size = 1) for p in paths if p is not None]
with (
patch("huggingface_hub.list_repo_files", lambda *_a, **_k: [MAIN]),
patch("huggingface_hub.get_paths_info", fake_get_paths_info),
patch("huggingface_hub.try_to_load_from_cache", lambda *_a, **_k: None),
patch("core.inference.llama_cpp.hf_hub_download_with_xet_fallback", fake_download),
):
out = backend._download_gguf(hf_repo = REPO, hf_variant = VARIANT)
assert downloaded == [MAIN]
assert out == f"/fake/{REPO}/{MAIN}"
def test_force_redownloads_despite_cache(self, hf_cache):
"""A forced download ignores a complete cached copy."""
backend = LlamaCppBackend()
_build_cache(hf_cache, REPO, {MAIN: 4})
downloaded: list[str] = []
def fake_download(
repo_id,
filename,
token = None,
**kwargs,
):
assert kwargs.get("force_download") is True
downloaded.append(filename)
return f"/fake/{repo_id}/{filename}"
def fake_get_paths_info(
_repo_id,
paths,
token = None,
):
return [_types.SimpleNamespace(path = p, size = 1) for p in paths if p is not None]
with (
patch("huggingface_hub.list_repo_files", lambda *_a, **_k: [MAIN]),
patch("huggingface_hub.get_paths_info", fake_get_paths_info),
patch("huggingface_hub.try_to_load_from_cache", lambda *_a, **_k: None),
patch("core.inference.llama_cpp.hf_hub_download_with_xet_fallback", fake_download),
):
out = backend._download_gguf(hf_repo = REPO, hf_variant = VARIANT, force = True)
assert downloaded == [MAIN]
assert out == f"/fake/{REPO}/{MAIN}"
def test_split_reused_only_when_colocated(self, hf_cache):
backend = LlamaCppBackend()
shard1 = f"gemma-test-{VARIANT}-00001-of-00002.gguf"
shard2 = f"gemma-test-{VARIANT}-00002-of-00002.gguf"
snap = _build_cache(hf_cache, REPO, {shard1: 4, shard2: 4})
with (
patch("huggingface_hub.list_repo_files", lambda *_a, **_k: [shard1, shard2]),
patch("huggingface_hub.get_paths_info", _fail_get_paths_info),
patch("core.inference.llama_cpp.hf_hub_download_with_xet_fallback", _fail_download),
):
out = backend._download_gguf(hf_repo = REPO, hf_variant = VARIANT)
assert out == str(snap / shard1)
def test_partial_split_set_downloads(self, hf_cache):
"""A partial split set is not reused."""
backend = LlamaCppBackend()
shard1 = f"gemma-test-{VARIANT}-00001-of-00002.gguf"
shard2 = f"gemma-test-{VARIANT}-00002-of-00002.gguf"
_build_cache(hf_cache, REPO, {shard1: 4})
downloaded: list[str] = []
def fake_download(
repo_id,
filename,
token = None,
**_kwargs,
):
downloaded.append(filename)
return f"/fake/{repo_id}/{filename}"
def fake_get_paths_info(
_repo_id,
paths,
token = None,
):
return [_types.SimpleNamespace(path = p, size = 4) for p in paths if p is not None]
with (
patch("huggingface_hub.list_repo_files", lambda *_a, **_k: [shard1, shard2]),
patch("huggingface_hub.get_paths_info", fake_get_paths_info),
patch("huggingface_hub.try_to_load_from_cache", lambda *_a, **_k: None),
patch("core.inference.llama_cpp.hf_hub_download_with_xet_fallback", fake_download),
):
out = backend._download_gguf(hf_repo = REPO, hf_variant = VARIANT)
assert downloaded == [shard1, shard2]
assert out == f"/fake/{REPO}/{shard1}"
def test_reuse_prefers_newest_snapshot_after_update(self, hf_cache):
"""Loads prefer the newest complete snapshot."""
import os
backend = LlamaCppBackend()
old_snap = _build_cache(hf_cache, REPO, {MAIN: 4}, snapshot_sha = "a" * 40)
new_snap = _build_cache(hf_cache, REPO, {MAIN: 6}, snapshot_sha = "b" * 40)
os.utime(old_snap, (1_000_000, 1_000_000))
os.utime(new_snap, (2_000_000, 2_000_000))
with (
patch("huggingface_hub.list_repo_files", lambda *_a, **_k: [MAIN]),
patch("huggingface_hub.get_paths_info", _fail_get_paths_info),
patch("core.inference.llama_cpp.hf_hub_download_with_xet_fallback", _fail_download),
):
out = backend._download_gguf(hf_repo = REPO, hf_variant = VARIANT)
assert out == str(new_snap / MAIN)
def test_low_disk_fallback_reuses_cached_copy(self, hf_cache):
backend = LlamaCppBackend()
fallback = "gemma-test-Q2_K.gguf"
snap = _build_cache(hf_cache, REPO, {fallback: 4})
def fake_get_paths_info(
_repo,
paths,
*,
revision = None,
token = None,
):
size = 4 if revision == snap.name else 100
return [_types.SimpleNamespace(path = path, size = size) for path in paths]
with (
patch("huggingface_hub.list_repo_files", lambda *_a, **_k: [MAIN]),
patch("huggingface_hub.get_paths_info", fake_get_paths_info),
patch("huggingface_hub.try_to_load_from_cache", lambda *_a, **_k: None),
patch("shutil.disk_usage", lambda *_a, **_k: _types.SimpleNamespace(free = 10)),
patch.object(
backend,
"_find_smallest_fitting_variant",
lambda *_a, **_k: (fallback, 4, []),
),
patch("core.inference.llama_cpp.hf_hub_download_with_xet_fallback", _fail_download),
):
out = backend._download_gguf(hf_repo = REPO, hf_variant = VARIANT)
assert out == str(snap / fallback)
def test_companion_prefers_main_snapshot_sibling(self, hf_cache):
"""A cached mmproj is reused from the main model's snapshot."""
backend = LlamaCppBackend()
snap = _build_cache(hf_cache, REPO, {MAIN: 4, "mmproj-F16.gguf": 2})
def _fail_list(*_args, **_kwargs):
raise AssertionError("snapshot sibling must resolve without a repo listing")
with patch("huggingface_hub.list_repo_files", _fail_list):
out = backend._download_mmproj(hf_repo = REPO, near_path = str(snap / MAIN))
assert out == str(snap / "mmproj-F16.gguf")
def test_companion_finds_snapshot_through_hf_symlink(self, hf_cache):
backend = LlamaCppBackend()
snap = _build_cache(hf_cache, REPO, {})
blobs = snap.parent.parent / "blobs"
main_blob = blobs / "main"
mmproj_blob = blobs / "mmproj"
main_blob.write_bytes(b"main")
mmproj_blob.write_bytes(b"mmproj")
try:
(snap / MAIN).symlink_to(main_blob)
(snap / "mmproj-F16.gguf").symlink_to(mmproj_blob)
except OSError as exc:
pytest.skip(f"symlinks unavailable: {exc}")
with patch("huggingface_hub.list_repo_files", _fail_download):
out = backend._download_mmproj(hf_repo = REPO, near_path = str(snap / MAIN))
assert out == str(snap / "mmproj-F16.gguf")
def test_companion_does_not_download_during_hub_job(self, hf_cache):
backend = LlamaCppBackend()
snap = _build_cache(hf_cache, REPO, {MAIN: 4})
registry = _types.SimpleNamespace(active_job_refs = lambda _repo: [object()])
with (
patch("huggingface_hub.list_repo_files", _fail_download),
patch("hub.utils.download_registry.get_models_registry", lambda: registry),
patch("core.inference.llama_cpp.hf_hub_download_with_xet_fallback", _fail_download),
):
out = backend._download_mmproj(hf_repo = REPO, near_path = str(snap / MAIN))
assert out is None
class TestCachedGgufForLoadProbe:
def test_complete_copy_found(self, hf_cache):
snap = _build_cache(hf_cache, REPO, {MAIN: 4})
assert cached_gguf_for_load(REPO, VARIANT) == str(snap / MAIN)
def test_absent_copy_is_none(self, hf_cache):
assert cached_gguf_for_load(REPO, VARIANT) is None
def test_partial_split_is_none(self, hf_cache):
shard1 = f"gemma-test-{VARIANT}-00001-of-00002.gguf"
_build_cache(hf_cache, REPO, {shard1: 4})
assert cached_gguf_for_load(REPO, VARIANT) is None
def test_partial_new_snapshot_does_not_hide_complete_split(self, hf_cache):
import os
shard1 = f"gemma-test-{VARIANT}-00001-of-00002.gguf"
shard2 = f"gemma-test-{VARIANT}-00002-of-00002.gguf"
old = _build_cache(
hf_cache,
REPO,
{shard1: 4, shard2: 4},
snapshot_sha = "a" * 40,
)
new = _build_cache(hf_cache, REPO, {shard1: 4}, snapshot_sha = "b" * 40)
os.utime(old, (1_000_000, 1_000_000))
os.utime(new, (2_000_000, 2_000_000))
assert cached_gguf_for_load(REPO, VARIANT) == str(old / shard1)
def test_split_requires_every_declared_shard(self, hf_cache):
shard1 = f"gemma-test-{VARIANT}-00001-of-00003.gguf"
shard2 = f"gemma-test-{VARIANT}-00002-of-00003.gguf"
_build_cache(hf_cache, REPO, {shard1: 4, shard2: 4})
assert cached_gguf_for_load(REPO, VARIANT) is None
def test_required_mmproj_must_share_main_snapshot(self, hf_cache):
snap = _build_cache(hf_cache, REPO, {MAIN: 4})
assert cached_gguf_for_load(REPO, VARIANT) == str(snap / MAIN)
assert cached_gguf_for_load(REPO, VARIANT, require_mmproj = True) is None
(snap / "mmproj-F16.gguf").write_bytes(b"mmproj")
assert cached_gguf_for_load(REPO, VARIANT, require_mmproj = True) == str(snap / MAIN)
def test_required_mmproj_scans_past_newer_main_only_snapshot(self, hf_cache):
import os
old = _build_cache(
hf_cache,
REPO,
{MAIN: 4, "mmproj-F16.gguf": 2},
snapshot_sha = "a" * 40,
)
new = _build_cache(hf_cache, REPO, {MAIN: 4}, snapshot_sha = "b" * 40)
os.utime(old, (1_000_000, 1_000_000))
os.utime(new, (2_000_000, 2_000_000))
assert cached_gguf_for_load(REPO, VARIANT, require_mmproj = True) == str(old / MAIN)
class TestLoadHubDownloadExclusion:
def test_in_flight_marker_counts_and_normalizes_case(self):
assert not hf_gguf_load_in_flight(REPO)
with gguf_load_in_flight(REPO):
assert hf_gguf_load_in_flight(REPO.upper())
with gguf_load_in_flight(REPO.lower()):
assert hf_gguf_load_in_flight(REPO)
assert hf_gguf_load_in_flight(REPO)
assert not hf_gguf_load_in_flight(REPO)
def test_marker_noops_for_local_loads(self):
with gguf_load_in_flight(None):
assert not hf_gguf_load_in_flight("")
def test_marker_cleared_on_exception(self):
with pytest.raises(RuntimeError):
with gguf_load_in_flight(REPO):
raise RuntimeError("boom")
assert not hf_gguf_load_in_flight(REPO)
def test_hub_download_refused_while_load_in_flight(self):
from fastapi import HTTPException
from hub.schemas.downloads import DownloadModelRequest
from hub.services.models import downloads as dl
body = DownloadModelRequest(repo_id = REPO, gguf_variant = VARIANT)
with (
patch.object(dl, "resolve_cached_repo_id_case", lambda repo_id, repo_type: repo_id),
gguf_load_in_flight(REPO),
):
with pytest.raises(HTTPException) as exc_info:
asyncio.run(dl.download_model_response(body))
assert exc_info.value.status_code == 409
assert "load" in exc_info.value.detail.lower()
def test_hub_download_rechecks_marker_before_claim(self):
from fastapi import HTTPException
from hub.schemas.downloads import DownloadModelRequest
from hub.services.models import downloads as dl
scope = None
def mark_load(*_args, **_kwargs):
nonlocal scope
if scope is None:
scope = gguf_load_in_flight(REPO)
scope.__enter__()
return frozenset()
class _Registry:
def claim(self, *_args, admission_check, **_kwargs):
assert admission_check() is False
return False, "admission_blocked"
def current_generation(self, _key):
return 0
registry = _Registry()
body = DownloadModelRequest(repo_id = REPO, gguf_variant = VARIANT)
try:
with (
patch.object(dl, "resolve_cached_repo_id_case", lambda repo_id, repo_type: repo_id),
patch.object(dl.gguf_variants, "gguf_variant_blob_hashes", mark_load),
patch.object(dl, "_registry", registry),
):
with pytest.raises(HTTPException) as exc_info:
asyncio.run(dl.download_model_response(body))
finally:
if scope is not None:
scope.__exit__(None, None, None)
assert exc_info.value.status_code == 409
def test_registry_admission_check_prevents_claim(self):
from hub.utils.download_registry import DownloadRegistry, TRANSPORT_HTTP
registry = DownloadRegistry()
claimed, state = registry.claim(
f"{REPO}::{VARIANT}",
TRANSPORT_HTTP,
repo_type = "model",
repo_id = REPO,
variant = VARIANT,
admission_check = lambda: False,
)
assert claimed is False
assert state == "admission_blocked"
assert registry.active_jobs(REPO) == {}
def test_same_variant_job_stays_visible_during_retry_handoff(self):
from hub.utils.download_registry import DownloadRegistry, TRANSPORT_XET
from core.inference.llama_cpp import _hub_download_blocks_gguf_load
registry = DownloadRegistry()
key = f"{REPO}::{VARIANT}"
claimed, _ = registry.claim(
key,
TRANSPORT_XET,
repo_type = "model",
repo_id = REPO,
variant = VARIANT,
)
assert claimed is True
assert registry.has_active_variant(REPO, VARIANT.lower()) is True
registry.release_active_slot(key)
assert registry.active_jobs(REPO) == {}
assert registry.active_job_refs(REPO)
assert registry.has_active_variant(REPO, VARIANT) is True
with (
patch("hub.utils.download_registry.get_models_registry", lambda: registry),
patch(
"core.inference.llama_cpp.cached_gguf_for_load",
side_effect = AssertionError("same-variant jobs must block before cache reuse"),
),
):
assert _hub_download_blocks_gguf_load(REPO, VARIANT) is True
registry.set_job(key, "complete")
assert registry.has_active_variant(REPO, VARIANT) is False
def test_other_variant_job_still_allows_complete_cached_load(self):
from core.inference.llama_cpp import _hub_download_blocks_gguf_load
from hub.utils.download_registry import DownloadRegistry, TRANSPORT_HTTP
registry = DownloadRegistry()
registry.claim(
f"{REPO}::Q8_0",
TRANSPORT_HTTP,
repo_type = "model",
repo_id = REPO,
variant = "Q8_0",
)
with (
patch("hub.utils.download_registry.get_models_registry", lambda: registry),
patch(
"core.inference.llama_cpp.cached_gguf_for_load",
return_value = "/cached/model.gguf",
) as cached_probe,
):
assert _hub_download_blocks_gguf_load(REPO, VARIANT) is False
cached_probe.assert_called_once_with(
REPO,
VARIANT,
require_mmproj = False,
verify_sizes = True,
hf_token = None,
)
def test_cancelled_request_keeps_marker_until_load_thread_finishes(self):
from core.inference.llama_cpp import _with_gguf_load_marker
started = threading.Event()
release = threading.Event()
finished = threading.Event()
class FakeBackend:
@_with_gguf_load_marker
def load_model(self, *, hf_repo):
started.set()
release.wait(timeout = 2)
finished.set()
return True
async def scenario():
with patch(
"core.inference.llama_cpp._hub_download_blocks_gguf_load",
return_value = False,
):
task = asyncio.create_task(
asyncio.to_thread(FakeBackend().load_model, hf_repo = REPO)
)
assert await asyncio.to_thread(started.wait, 1)
task.cancel()
with pytest.raises(asyncio.CancelledError):
await task
assert hf_gguf_load_in_flight(REPO)
release.set()
assert await asyncio.to_thread(finished.wait, 1)
for _ in range(100):
if not hf_gguf_load_in_flight(REPO):
break
await asyncio.sleep(0.001)
assert not hf_gguf_load_in_flight(REPO)
asyncio.run(scenario())
def test_load_marker_precedes_hub_guard_and_unload(self):
source = (Path(__file__).resolve().parent.parent / "routes" / "inference.py").read_text()
# _load_model_impl has more than one `if config.is_gguf:`, so anchor on
# the branch that actually owns the load marker rather than the first
# one in the file, which belongs to an earlier check.
marker = source.index("enter_context(gguf_load_in_flight")
gguf_branch_start = source.rindex("if config.is_gguf:", 0, marker)
gguf_branch = source[gguf_branch_start:]
# The gguf_load_in_flight marker must be entered before the hub-download
# guard and the unload so a concurrent load can't race the download
# manager. The llama_extra_args inheritance moved out of the branch into
# _resolve_inherited_extra_args, which must still run BEFORE it: the
# inherited value (e.g. a carried --no-mmproj) shapes the guard's
# require_mmproj. Anchor on the call form so the assertion pins the
# endpoint's call site, not the function definition.
assert source.index("= _resolve_inherited_extra_args(") < gguf_branch_start
assert (
gguf_branch.index("enter_context(gguf_load_in_flight")
< gguf_branch.index("_hub_download_blocks_gguf_load")
< gguf_branch.index("unsloth_backend.unload_model")
)
llama_source = (
Path(__file__).resolve().parent.parent / "core" / "inference" / "llama_cpp.py"
).read_text()
assert "@_with_gguf_load_marker\n def load_model(" in llama_source
def _capture_hub_guard_require_mmproj(
self,
stored_extra_args,
request_extra_args = None,
):
"""Drive /load's GGUF path and return the hub guard's require_mmproj.
The guard reports a conflicting download, so the 409 is the observation
point and no llama-server ever starts.
"""
import core.inference.llama_cpp as llama_cpp_module
from fastapi import HTTPException
from models.inference import LoadRequest
route = _load_route_module(
"inference_route_module_for_inherited_extra_args_test",
"routes/inference.py",
)
captured = {}
def _fake_blocks(
repo,
variant,
*,
require_mmproj,
hf_token = None,
):
captured["repo"] = repo
captured["variant"] = variant
captured["require_mmproj"] = require_mmproj
return True
# A vision GGUF: require_mmproj is True unless the extras say --no-mmproj.
config = SimpleNamespace(
is_gguf = True,
is_lora = False,
is_vision = True,
is_audio = False,
audio_type = None,
has_audio_input = False,
gguf_hf_repo = REPO,
gguf_variant = VARIANT,
gguf_file = None,
gguf_mmproj_file = None,
identifier = REPO,
display_name = REPO,
)
# Pass-through extras the running backend recorded for the last load.
llama_backend = SimpleNamespace(
is_loaded = False,
extra_args = list(stored_extra_args),
extra_args_source = (REPO, VARIANT),
hf_variant = VARIANT,
model_identifier = REPO,
)
request = LoadRequest(
model_path = REPO,
gguf_variant = VARIANT,
llama_extra_args = request_extra_args,
)
with (
patch.object(
route,
"ModelConfig",
SimpleNamespace(from_identifier = lambda **_kwargs: config),
),
patch.object(route, "get_llama_cpp_backend", lambda: llama_backend),
patch.object(
route,
"get_inference_backend",
lambda: SimpleNamespace(active_model_name = None),
),
patch.object(route, "_resolve_gguf_gpu_ids_for_request", _no_gguf_gpu_ids),
patch.object(route, "_guard_chat_load_against_training", return_value = None),
patch.object(route, "_effective_load_in_4bit", return_value = False),
patch.object(route, "_hf_offline_if_dns_dead", nullcontext),
patch.object(route.asyncio, "to_thread", new = _inline_to_thread),
patch.object(llama_cpp_module, "_hub_download_blocks_gguf_load", _fake_blocks),
):
with pytest.raises(HTTPException) as exc_info:
asyncio.run(
route._load_model_impl(
request,
SimpleNamespace(
app = SimpleNamespace(
state = SimpleNamespace(llama_parallel_slots = 1),
),
),
current_subject = "test-user",
)
)
assert exc_info.value.status_code == 409
assert captured["repo"] == REPO
return captured["require_mmproj"]
def test_inherited_extra_args_shape_hub_guard_require_mmproj(self):
# Inheritance must resolve before the hub-download guard: an inherited
# --no-mmproj decides require_mmproj, so resolving later rejects a load
# over a download the effective arguments disable (#7251).
assert self._capture_hub_guard_require_mmproj(["--no-mmproj"]) is False
# Control: nothing to inherit, so a vision GGUF still needs its mmproj.
assert self._capture_hub_guard_require_mmproj([]) is True
# An explicit request list wins over the stored one, both ways.
assert (
self._capture_hub_guard_require_mmproj([], request_extra_args = ["--no-mmproj"]) is False
)
assert (
self._capture_hub_guard_require_mmproj(["--no-mmproj"], request_extra_args = []) is True
)