diff --git a/backend/services/hf_cache_repair.py b/backend/services/hf_cache_repair.py index fa82f085..9fe30b17 100644 --- a/backend/services/hf_cache_repair.py +++ b/backend/services/hf_cache_repair.py @@ -238,9 +238,11 @@ def repair_repo_cache(repo_id: str, cache_dir: str | None = None) -> dict: from huggingface_hub import snapshot_download - dl_kwargs: dict = {"repo_id": repo_id, "revision": revision} - if cache_dir: - dl_kwargs["cache_dir"] = cache_dir + dl_kwargs: dict = { + "repo_id": repo_id, + "revision": revision, + "cache_dir": cache_root, + } endpoint = os.environ.get("HF_ENDPOINT") if endpoint: dl_kwargs["endpoint"] = endpoint diff --git a/backend/services/model_manager.py b/backend/services/model_manager.py index 82ff30f3..2c432b6a 100644 --- a/backend/services/model_manager.py +++ b/backend/services/model_manager.py @@ -1665,12 +1665,17 @@ def _repair_model_cache(checkpoint: str, *, force: bool = False) -> bool: try: from services.hf_cache_repair import hf_cache_home from services.hf_revisions import installed_revision - revision = installed_revision(checkpoint, hf_cache_home()) + cache_root = hf_cache_home() + revision = installed_revision(checkpoint, cache_root) except (OSError, ValueError) as revision_err: _last_repair_error = str(revision_err) logger.warning("Refusing unpinned model repair for %s: %s", checkpoint, revision_err) return False - dl_kwargs: dict = {"repo_id": checkpoint, "revision": revision} + dl_kwargs: dict = { + "repo_id": checkpoint, + "revision": revision, + "cache_dir": cache_root, + } # Explicit endpoint (HF_ENDPOINT / pref) wins; otherwise the automatic # endpoint selection's cached pick applies (services.endpoint_race). try: diff --git a/tests/test_hf_cache_repair.py b/tests/test_hf_cache_repair.py index 8f3c0149..2e271852 100644 --- a/tests/test_hf_cache_repair.py +++ b/tests/test_hf_cache_repair.py @@ -148,11 +148,11 @@ def test_repair_removes_only_broken_and_redownloads(tmp_path, monkeypatch): _symlink_or_skip(os.path.join("..", "..", "blobs", "MISSING"), str(snap / "model.safetensors")) calls = [] + monkeypatch.setenv("HF_HUB_CACHE", str(cache)) monkeypatch.setattr(huggingface_hub, "snapshot_download", lambda **k: calls.append(k)) - summary = hf_cache_repair.repair_repo_cache("test/checkpoint", - cache_dir=str(cache)) + summary = hf_cache_repair.repair_repo_cache("test/checkpoint") assert summary["found"] == 1 assert summary["removed"] == 1 assert summary["restored"] is True diff --git a/tests/test_model_cache_repair.py b/tests/test_model_cache_repair.py index 807cc471..08605370 100644 --- a/tests/test_model_cache_repair.py +++ b/tests/test_model_cache_repair.py @@ -191,7 +191,7 @@ def test_repair_skipped_in_offline_mode(model_manager, monkeypatch): assert called == [] # no download attempted offline -def test_repair_invokes_snapshot_download(model_manager, monkeypatch): +def test_repair_invokes_snapshot_download(model_manager, monkeypatch, tmp_path): """Repair re-fetches the repo via snapshot_download (resume/fill missing).""" calls = [] @@ -200,11 +200,13 @@ def test_repair_invokes_snapshot_download(model_manager, monkeypatch): return "/cache/test/checkpoint" import huggingface_hub + monkeypatch.setenv("HF_HUB_CACHE", str(tmp_path)) monkeypatch.setattr(huggingface_hub, "snapshot_download", fake_snapshot_download) assert model_manager._repair_model_cache("test/checkpoint") is True assert calls and calls[0]["repo_id"] == "test/checkpoint" assert calls[0]["revision"] == "a" * 40 + assert calls[0]["cache_dir"] == str(tmp_path) def test_repair_returns_false_when_download_fails(model_manager, monkeypatch):