diff --git a/studio/backend/core/data_recipe/jobs/manager.py b/studio/backend/core/data_recipe/jobs/manager.py index 0e0044702e..135c9fccf6 100644 --- a/studio/backend/core/data_recipe/jobs/manager.py +++ b/studio/backend/core/data_recipe/jobs/manager.py @@ -27,7 +27,6 @@ from .constants import ( ) from .parse import apply_update, coerce_event, parse_log_message from .types import Job -from .worker import run_job_process from loggers import get_logger logger = get_logger(__name__) @@ -169,12 +168,18 @@ class JobManager: native_path_secret_removed_for_child_start, run_without_native_path_secret, ) + from utils.hf_cache_settings import child_environment_for_spawn, get_hf_cache_paths - with native_path_secret_removed_for_child_start(): + cache_env = get_hf_cache_paths().child_env({}) + + with ( + child_environment_for_spawn(cache_env), + native_path_secret_removed_for_child_start(), + ): mp_q = _CTX.Queue() proc = _CTX.Process( target = run_without_native_path_secret, - args = (run_job_process,), + args = ("core.data_recipe.jobs.worker", "run_job_process", cache_env), kwargs = {"event_queue": mp_q, "recipe": recipe, "run": run_payload}, daemon = True, ) diff --git a/studio/backend/core/export/orchestrator.py b/studio/backend/core/export/orchestrator.py index 6d1a928f2e..aaf48615f0 100644 --- a/studio/backend/core/export/orchestrator.py +++ b/studio/backend/core/export/orchestrator.py @@ -230,16 +230,20 @@ class ExportOrchestrator: native_path_secret_removed_for_child_start, run_without_native_path_secret, ) + from utils.hf_cache_settings import child_environment_for_spawn, get_hf_cache_paths - from .worker import run_export_process + cache_env = get_hf_cache_paths().child_env({}) - with native_path_secret_removed_for_child_start(): + with ( + child_environment_for_spawn(cache_env), + native_path_secret_removed_for_child_start(), + ): self._cmd_queue = _CTX.Queue() self._resp_queue = _CTX.Queue() self._proc = _CTX.Process( target = run_without_native_path_secret, - args = (run_export_process,), + args = ("core.export.worker", "run_export_process", cache_env), kwargs = { "cmd_queue": self._cmd_queue, "resp_queue": self._resp_queue, diff --git a/studio/backend/core/inference/audio_codecs.py b/studio/backend/core/inference/audio_codecs.py index 93c7da72cb..b59f2bcce0 100644 --- a/studio/backend/core/inference/audio_codecs.py +++ b/studio/backend/core/inference/audio_codecs.py @@ -76,8 +76,14 @@ class AudioCodecManager: if self._snac_model is not None: return from snac import SNAC + from utils.hf_cache_settings import active_hf_hub_cache - self._snac_model = SNAC.from_pretrained("hubertsiuzdak/snac_24khz").to(device).eval() + # Route weights to the selected cache; this can run in the main process. + self._snac_model = ( + SNAC.from_pretrained("hubertsiuzdak/snac_24khz", cache_dir = active_hf_hub_cache()) + .to(device) + .eval() + ) logger.info("Loaded SNAC codec (24kHz)") def _load_bicodec( diff --git a/studio/backend/core/inference/llama_cpp.py b/studio/backend/core/inference/llama_cpp.py index 1c9c76ebe9..fc58442a79 100644 --- a/studio/backend/core/inference/llama_cpp.py +++ b/studio/backend/core/inference/llama_cpp.py @@ -579,7 +579,14 @@ def _swa_entry_from_layer_types(lt) -> Optional[object]: def _fetch_swa_entry_from_hf(repo_id: str) -> Optional[object]: try: from huggingface_hub import hf_hub_download - cfg_path = hf_hub_download(repo_id, "config.json", repo_type = "model") + from utils.hf_cache_settings import active_hf_hub_cache + + cfg_path = hf_hub_download( + repo_id, + "config.json", + repo_type = "model", + cache_dir = active_hf_hub_cache(), + ) with open(cfg_path) as f: cfg = json.load(f) except Exception: @@ -981,6 +988,7 @@ def _cached_hf_snapshot_file( filename: str, *, expected_size: Optional[int] = None, + cache_dir: Optional[str] = None, ) -> Optional[str]: """Return a cached snapshot file even when HF's current-ref probe misses it.""" if not filename: @@ -989,8 +997,22 @@ def _cached_hf_snapshot_file( if not parts or any(part in (".", "..") for part in parts): return None try: - from utils.models.model_config import _iter_hf_cache_snapshots - for snap in _iter_hf_cache_snapshots(repo_id): + if cache_dir is None: + from utils.models.model_config import _iter_hf_cache_snapshots + snapshots = _iter_hf_cache_snapshots(repo_id) + else: + from hub.utils.hf_cache_state import iter_active_repo_cache_dirs + snapshots = ( + snapshot + for repo_dir in iter_active_repo_cache_dirs( + "model", + repo_id, + root = Path(cache_dir), + ) + for snapshot in (repo_dir / "snapshots").glob("*") + if snapshot.is_dir() + ) + for snap in snapshots: candidate = snap.joinpath(*parts) if not candidate.is_file(): continue @@ -1232,6 +1254,16 @@ def _snapshot_dir_of(path: str) -> Optional[Path]: return None +def _hub_cache_dir_for_snapshot_path(path: Optional[str]) -> Optional[str]: + """Return the HF Hub cache root that owns a snapshot-contained path.""" + if not path: + return None + snapshot = _snapshot_dir_of(path) + if snapshot is None or snapshot.parent.name != "snapshots": + return None + return str(snapshot.parent.parent.parent) + + def _companion_snapshot_sibling( near_path: str, pick: Callable[[list[str]], Optional[str]] ) -> Optional[str]: @@ -5070,6 +5102,9 @@ class LlamaCppBackend: touching the shared one; defaults to the shared event. """ cancel_event = cancel_event if cancel_event is not None else self._cancel_event + from utils.hf_cache_settings import get_hf_cache_paths + + download_cache_dir = str(get_hf_cache_paths().hub_cache) try: import huggingface_hub # noqa: F401 -- presence check only except ImportError: @@ -5165,7 +5200,11 @@ class LlamaCppBackend: if not p.size: continue try: - cached_path = try_to_load_from_cache(hf_repo, p.path) + cached_path = try_to_load_from_cache( + hf_repo, + p.path, + cache_dir = download_cache_dir, + ) except Exception: cached_path = None if ( @@ -5176,6 +5215,7 @@ class LlamaCppBackend: hf_repo, p.path, expected_size = p.size, + cache_dir = download_cache_dir, ) if isinstance(cached_path, str) and os.path.exists(cached_path): try: @@ -5189,12 +5229,8 @@ class LlamaCppBackend: total_download_bytes = max(0, total_bytes - already_cached_bytes) if total_download_bytes > 0: - cache_dir = os.environ.get( - "HF_HUB_CACHE", - str(Path.home() / ".cache" / "huggingface" / "hub"), - ) - Path(cache_dir).mkdir(parents = True, exist_ok = True) - free_bytes = shutil.disk_usage(cache_dir).free + Path(download_cache_dir).mkdir(parents = True, exist_ok = True) + free_bytes = shutil.disk_usage(download_cache_dir).free total_gb = total_download_bytes / (1024**3) free_gb = free_bytes / (1024**3) @@ -5212,7 +5248,7 @@ class LlamaCppBackend: # surface the disk shortfall for the requested variant. raise RuntimeError( f"Not enough disk space to download {gguf_filename}. " - f"Only {free_gb:.1f} GB free in {cache_dir}" + f"Only {free_gb:.1f} GB free in {download_cache_dir}" ) smaller = self._find_smallest_fitting_variant( hf_repo, @@ -5243,7 +5279,7 @@ class LlamaCppBackend: else: raise RuntimeError( f"Not enough disk space to download any variant. " - f"Only {free_gb:.1f} GB free in {cache_dir}" + f"Only {free_gb:.1f} GB free in {download_cache_dir}" ) except RuntimeError: raise @@ -5266,6 +5302,7 @@ class LlamaCppBackend: cancel_event = cancel_event, on_status = lambda m: logger.info(m), force_download = force, + cache_dir = download_cache_dir, ) for shard in gguf_extra_shards: if cancel_event.is_set(): @@ -5277,6 +5314,7 @@ class LlamaCppBackend: hf_token, cancel_event = cancel_event, force_download = force, + cache_dir = download_cache_dir, ) except Exception as e: if isinstance(e, RuntimeError) and "Cancelled" in str(e): @@ -5322,6 +5360,12 @@ class LlamaCppBackend: logger.info("Reusing cached %s: %s", label, cached) return cached + from utils.hf_cache_settings import get_hf_cache_paths + + companion_cache_dir = _hub_cache_dir_for_snapshot_path(near_path) or str( + get_hf_cache_paths().hub_cache + ) + if _hub_download_in_flight(hf_repo): logger.info("Skipping %s download while a hub download is active", label) return None @@ -5356,7 +5400,7 @@ class LlamaCppBackend: if target is None: try: from utils.models.model_config import _iter_hf_cache_snapshots - for snap in _iter_hf_cache_snapshots(hf_repo): + for snap in _iter_hf_cache_snapshots(hf_repo, companion_cache_dir): rel_files = _gguf_snapshot_files(snap) target = pick(rel_files) if target is not None: @@ -5374,7 +5418,11 @@ class LlamaCppBackend: # hf_hub_download with hf_repo would miss the canonical file and silently # drop the companion. _cached_hf_snapshot_file scans every case variant. if _hf_env_offline(): - cached = _cached_hf_snapshot_file(hf_repo, target) + cached = _cached_hf_snapshot_file( + hf_repo, + target, + cache_dir = companion_cache_dir, + ) if cached: logger.info("Resolved %s from local HF cache: %s", label, cached) return cached @@ -5387,6 +5435,7 @@ class LlamaCppBackend: target, hf_token, cancel_event = cancel_event, + cache_dir = companion_cache_dir, ) except Exception as e: logger.warning(f"Could not download {label}: {e}") @@ -5417,7 +5466,12 @@ class LlamaCppBackend: near_path = near_path, ) - def _cached_repo_mtp_drafter(self, hf_repo: str) -> Optional[str]: + def _cached_repo_mtp_drafter( + self, + hf_repo: str, + *, + cache_dir: Optional[str] = None, + ) -> Optional[str]: """A drafter already in this repo's local HF cache, reused offline when a fresh copy can't be fetched. Prefers a repo-root ``mtp-*.gguf`` across all cached snapshots; else an existing ``MTP/`` copy (any precision -- the @@ -5427,7 +5481,12 @@ class LlamaCppBackend: roots: list[Path] = [] subdirs: list[Path] = [] - for snap in _iter_hf_cache_snapshots(hf_repo): # newest first + snapshots = ( + _iter_hf_cache_snapshots(hf_repo) + if cache_dir is None + else _iter_hf_cache_snapshots(hf_repo, cache_dir) + ) + for snap in snapshots: # newest first for f in sorted(_gguf_snapshot_files(snap)): if _is_companion_gguf_path(f) and "mmproj" not in f.lower(): (roots if "/" not in f else subdirs).append(snap / f) @@ -5480,7 +5539,10 @@ class LlamaCppBackend: # current cached file and refetch a changed one, so skip the probe here # rather than pair new weights with a stale draft. if _hf_env_offline(): - cached = self._cached_repo_mtp_drafter(hf_repo) + cached = self._cached_repo_mtp_drafter( + hf_repo, + cache_dir = _hub_cache_dir_for_snapshot_path(near_path), + ) if cached: logger.info(f"Reusing cached MTP drafter (offline): {cached}") return cached diff --git a/studio/backend/core/inference/local_model_resolver.py b/studio/backend/core/inference/local_model_resolver.py index 64ab38ec75..9e3eaeda3f 100644 --- a/studio/backend/core/inference/local_model_resolver.py +++ b/studio/backend/core/inference/local_model_resolver.py @@ -146,6 +146,7 @@ def _build_index() -> dict[str, _LocalGgufEntry]: _is_hidden_model, ) from utils.paths import legacy_hf_cache_dir, hf_default_cache_dir, lmstudio_model_dirs + from utils.hf_cache_settings import known_hf_hub_caches index: dict[str, _LocalGgufEntry] = {} seen_hf: set[str] = set() @@ -174,7 +175,12 @@ def _build_index() -> dict[str, _LocalGgufEntry]: except Exception as exc: logger.debug("auto-switch: ./models scan failed: %s", exc) try: - for hf_dir in (_resolve_hf_cache_dir(), legacy_hf_cache_dir(), hf_default_cache_dir()): + for hf_dir in ( + *known_hf_hub_caches(), + _resolve_hf_cache_dir(), + legacy_hf_cache_dir(), + hf_default_cache_dir(), + ): found += _scan_hf_once(hf_dir) except Exception as exc: logger.debug("auto-switch: HF cache scan failed: %s", exc) diff --git a/studio/backend/core/inference/orchestrator.py b/studio/backend/core/inference/orchestrator.py index 75ef9c2399..409132d605 100644 --- a/studio/backend/core/inference/orchestrator.py +++ b/studio/backend/core/inference/orchestrator.py @@ -217,10 +217,14 @@ class InferenceOrchestrator: native_path_secret_removed_for_child_start, run_without_native_path_secret, ) + from utils.hf_cache_settings import child_environment_for_spawn, get_hf_cache_paths - from .worker import run_inference_process + cache_env = get_hf_cache_paths().child_env({}) - with native_path_secret_removed_for_child_start(): + with ( + child_environment_for_spawn(cache_env), + native_path_secret_removed_for_child_start(), + ): self._cmd_queue = _CTX.Queue() self._resp_queue = _CTX.Queue() self._cancel_event = _CTX.Event() @@ -228,7 +232,7 @@ class InferenceOrchestrator: self._proc = _CTX.Process( target = run_without_native_path_secret, - args = (run_inference_process,), + args = ("core.inference.worker", "run_inference_process", cache_env), kwargs = { "cmd_queue": self._cmd_queue, "resp_queue": self._resp_queue, diff --git a/studio/backend/core/rag/embed_llama_server.py b/studio/backend/core/rag/embed_llama_server.py index b141e59422..facd989b27 100644 --- a/studio/backend/core/rag/embed_llama_server.py +++ b/studio/backend/core/rag/embed_llama_server.py @@ -188,7 +188,14 @@ class LlamaServerBackend: match = [f for f in files if variant in f.lower()] or files filename = sorted(match, key = len)[0] logger.info("resolving GGUF embedder %s/%s", repo, filename) - self._model_path = hf_hub_download(repo_id = repo, filename = filename, token = token) + from utils.hf_cache_settings import active_hf_hub_cache + + self._model_path = hf_hub_download( + repo_id = repo, + filename = filename, + token = token, + cache_dir = active_hf_hub_cache(), + ) self._model_repo = desired self._dim = None return self._model_path diff --git a/studio/backend/core/rag/embeddings.py b/studio/backend/core/rag/embeddings.py index 0c743e4ea4..3354585d2a 100644 --- a/studio/backend/core/rag/embeddings.py +++ b/studio/backend/core/rag/embeddings.py @@ -104,9 +104,15 @@ def _st_module_subdirs(name: str, token: str | None) -> tuple[str, ...]: else: from huggingface_hub import hf_hub_download from huggingface_hub.utils import EntryNotFoundError + from utils.hf_cache_settings import active_hf_hub_cache try: - local = hf_hub_download(name, "modules.json", token = token or None) + local = hf_hub_download( + name, + "modules.json", + token = token or None, + cache_dir = active_hf_hub_cache(), + ) except EntryNotFoundError: return () data = json.loads(open(local).read()) @@ -183,11 +189,16 @@ def _get(model_name: str | None = None): if _model is None or _name != name: _install_torchao_stub_once() from sentence_transformers import SentenceTransformer + from utils.hf_cache_settings import active_hf_hub_cache device = _device() logger.info("loading embedding model %s on %s", name, device) _guard_model_security(name, local_only) - st_kwargs = dict(device = device, model_kwargs = dtype_kwargs("float16")) + st_kwargs = dict( + device = device, + cache_folder = active_hf_hub_cache(), + model_kwargs = dtype_kwargs("float16"), + ) load_target = name if local_only: from utils.utils import hf_cache_snapshot_dir diff --git a/studio/backend/core/training/training.py b/studio/backend/core/training/training.py index 26d8c23b66..7cfa61d60f 100644 --- a/studio/backend/core/training/training.py +++ b/studio/backend/core/training/training.py @@ -929,16 +929,21 @@ class TrainingBackend: config["resolved_gpu_ids"] = resolved_gpu_ids config["gpu_selection"] = gpu_selection - from .worker import run_training_process + from utils.hf_cache_settings import child_environment_for_spawn, get_hf_cache_paths + + cache_env = get_hf_cache_paths().child_env({}) try: - with native_path_secret_removed_for_child_start(): + with ( + child_environment_for_spawn(cache_env), + native_path_secret_removed_for_child_start(), + ): event_queue = _CTX.Queue() stop_queue = _CTX.Queue() proc = _CTX.Process( target = run_without_native_path_secret, - args = (run_training_process,), + args = ("core.training.worker", "run_training_process", cache_env), kwargs = { "event_queue": event_queue, "stop_queue": stop_queue, @@ -991,6 +996,7 @@ class TrainingBackend: self._db_started_at = datetime.now(timezone.utc).isoformat() # Start each job Xet-first; keep config so a stall can respawn over HTTP. self._last_full_config = config + self._last_hf_cache_env = cache_env self._in_model_load = False self._xet_fallback_used = False self._needs_xet_respawn = False @@ -1400,7 +1406,11 @@ class TrainingBackend: self._last_full_config = config logger.warning("Respawning training worker with HF_HUB_DISABLE_XET=1 after Xet stall") - from .worker import run_training_process + cache_env = getattr(self, "_last_hf_cache_env", None) + if not cache_env: + from utils.hf_cache_settings import get_hf_cache_paths + cache_env = get_hf_cache_paths().child_env({}) + from utils.hf_cache_settings import child_environment_for_spawn # This run is active, so an install request 409s rather than proceeds: a reservation seen here # is transient (an aborting install or short lazy repair). Wait it out instead of stranding the @@ -1432,12 +1442,15 @@ class TrainingBackend: # crashed respawn cannot wedge is_training_active until restart. try: try: - with native_path_secret_removed_for_child_start(): + with ( + child_environment_for_spawn(cache_env), + native_path_secret_removed_for_child_start(), + ): event_queue = _CTX.Queue() stop_queue = _CTX.Queue() new_proc = _CTX.Process( target = run_without_native_path_secret, - args = (run_training_process,), + args = ("core.training.worker", "run_training_process", cache_env), kwargs = { "event_queue": event_queue, "stop_queue": stop_queue, diff --git a/studio/backend/hub/routes/datasets.py b/studio/backend/hub/routes/datasets.py index edf4f36ac0..7c7cc274d3 100644 --- a/studio/backend/hub/routes/datasets.py +++ b/studio/backend/hub/routes/datasets.py @@ -61,9 +61,11 @@ async def list_cached_datasets(current_subject: str = Depends(get_current_subjec @router.delete("/cached", response_model = DeleteCachedDatasetResponse) async def delete_cached_dataset( - repo_id: str = Body(..., embed = True), current_subject: str = Depends(get_current_subject) + repo_id: str = Body(..., embed = True), + cache_path: Optional[str] = Body(None, embed = True), + current_subject: str = Depends(get_current_subject), ): - return await cache_inventory.delete_cached_dataset_response(repo_id) + return await cache_inventory.delete_cached_dataset_response(repo_id, cache_path) @router.get("/download-progress", response_model = DownloadProgressResponse) diff --git a/studio/backend/hub/routes/inventory.py b/studio/backend/hub/routes/inventory.py index 1ffadf0544..dc3e3641bc 100644 --- a/studio/backend/hub/routes/inventory.py +++ b/studio/backend/hub/routes/inventory.py @@ -233,7 +233,8 @@ async def list_hidden_models(current_subject: str = Depends(get_current_subject) async def delete_cached_model( repo_id: str = Body(...), variant: Optional[str] = Body(None), + cache_path: Optional[str] = Body(None), hf_token: Optional[str] = Depends(get_hf_token), current_subject: str = Depends(get_current_subject), ): - return await deletion.delete_cached_model_response(repo_id, variant, hf_token) + return await deletion.delete_cached_model_response(repo_id, variant, hf_token, cache_path) diff --git a/studio/backend/hub/schemas/inventory.py b/studio/backend/hub/schemas/inventory.py index 19d6da3e11..ca0f4658a3 100644 --- a/studio/backend/hub/schemas/inventory.py +++ b/studio/backend/hub/schemas/inventory.py @@ -99,6 +99,10 @@ class LocalModelInfo(BaseModel): None, description = "HF repo id for cached models, e.g. org/model", ) + active_cache: Optional[bool] = Field( + None, + description = "Whether this HF entry belongs to the current download cache.", + ) base_model: Optional[str] = Field( None, description = "Base model from adapter_config.json when this is an adapter", diff --git a/studio/backend/hub/services/datasets/cache_inventory.py b/studio/backend/hub/services/datasets/cache_inventory.py index a180c9df58..c106b62ab7 100644 --- a/studio/backend/hub/services/datasets/cache_inventory.py +++ b/studio/backend/hub/services/datasets/cache_inventory.py @@ -20,12 +20,11 @@ from hub.utils import inventory_scan as hf_cache_scan from hub.utils.hf_cache_state import ( purge_partial_repo, purge_repo_cache_dirs, + resolve_delete_target_root, resolve_destructive_case_matches, ) from hub.utils.paths import ( - hf_default_cache_dir, is_valid_repo_id as _is_valid_repo_id, - legacy_hf_cache_dir, resolve_cached_repo_id_case, ) @@ -43,38 +42,8 @@ def _collect_hf_cache_scans() -> tuple[list, set[str]]: def _hf_hub_cache_roots() -> list[Path]: - roots: list[Path] = [] - seen: set[str] = set() - - def _add(path: Optional[Path]) -> None: - if path is None or not path.is_dir(): - return - try: - resolved = str(path.resolve()) - except OSError: - return - if resolved in seen: - return - seen.add(resolved) - roots.append(path) - - try: - from huggingface_hub.constants import HF_HUB_CACHE - _add(Path(HF_HUB_CACHE)) - except Exception: - pass - - hf_hub_cache = os.environ.get("HF_HUB_CACHE") - if hf_hub_cache: - _add(Path(hf_hub_cache).expanduser()) - - hf_home = os.environ.get("HF_HOME") - if hf_home: - _add(Path(hf_home).expanduser() / "hub") - - _add(legacy_hf_cache_dir()) - _add(hf_default_cache_dir()) - return roots + from hub.utils.hf_cache_state import hf_cache_roots + return hf_cache_roots() def _repo_id_from_hub_dataset_dir(name: str) -> str | None: @@ -207,6 +176,21 @@ def _repo_id_from_datasets_cache_dir(name: str) -> str | None: return repo_id if _is_valid_repo_id(repo_id) else None +def _is_processed_dataset_cache_path(repo_id: str, cache_path: str) -> bool: + """True when *cache_path* is this repo's processed Arrow cache dir + (``___`` directly under an HF_DATASETS_CACHE root). Such rows + have no Hub ``datasets--`` layout, so they are deleted via the processed + path and must not be rejected as an invalid cache_path.""" + try: + resolved = Path(cache_path).expanduser().resolve(strict = False) + except (OSError, RuntimeError, ValueError): + return False + if resolved.name.lower() != repo_id.replace("/", "___").lower(): + return False + roots = {r.resolve(strict = False) for r in _hf_datasets_cache_roots()} + return resolved.parent.resolve(strict = False) in roots + + def _processed_dataset_cache_size(path: Path) -> int: total = 0 try: @@ -361,7 +345,7 @@ async def list_cached_datasets_response() -> dict: ) from exc -async def delete_cached_dataset_response(repo_id: str) -> dict: +async def delete_cached_dataset_response(repo_id: str, cache_path: Optional[str] = None) -> dict: """Remove a cached dataset repo from the HF cache.""" if not _is_valid_repo_id(repo_id): raise HTTPException(status_code = 400, detail = "Invalid repo_id format") @@ -373,22 +357,40 @@ async def delete_cached_dataset_response(repo_id: str) -> dict: detail = "Cancel the active download before deleting.", ) try: - return await asyncio.to_thread(_delete_cached_dataset_blocking, repo_key) + return await asyncio.to_thread(_delete_cached_dataset_blocking, repo_key, cache_path) finally: downloads.registry.end_delete(repo_key) hf_cache_scan.invalidate_hf_cache_scans() -def _delete_cached_dataset_blocking(repo_id: str) -> dict: +def _delete_cached_dataset_blocking(repo_id: str, cache_path: Optional[str] = None) -> dict: scans, _seen_roots = _collect_hf_cache_scans() - candidate_entries = [] + # Group this dataset's copies by owning cache root, then target exactly one + # cache so a delete never removes copies in other, previously selected caches. + owners: dict = {} for hf_cache in scans: for repo_info in hf_cache.repos: if str(repo_info.repo_type) != "dataset": continue - if repo_info.repo_id.lower() == repo_id.lower(): - candidate_entries.append((hf_cache, repo_info)) + if repo_info.repo_id.lower() != repo_id.lower(): + continue + try: + owner = Path(repo_info.repo_path).parent.resolve(strict = False) + except (OSError, RuntimeError, ValueError): + continue + owners.setdefault(owner, []).append((hf_cache, repo_info)) + + target_root = resolve_delete_target_root("dataset", repo_id, cache_path, owners.keys()) + # A processed-only dataset row sends its Arrow cache path (___ + # under HF_DATASETS_CACHE), which is not a Hub datasets-- dir, so + # resolve_delete_target_root returns None. Accept it and fall through to the + # processed-cache delete rather than rejecting a legitimate row. + if target_root is None and not ( + cache_path and _is_processed_dataset_cache_path(repo_id, cache_path) + ): + raise HTTPException(status_code = 400, detail = "Invalid cache_path") + candidate_entries = owners.get(target_root, []) if target_root is not None else [] matched_repo_ids = resolve_destructive_repo_ids( repo_id, [str(repo_info.repo_id) for _hf_cache, repo_info in candidate_entries], @@ -414,7 +416,26 @@ def _delete_cached_dataset_blocking(repo_id: str) -> dict: exc_info = True, ) - processed_deleted, processed_failures = _delete_processed_dataset_cache(repo_id) + # Restrict the processed Arrow-cache delete to the selected cache's datasets + # root so it never removes copies under other cache homes. A processed + # cache_path scopes to its own root; a Hub target scopes to the datasets root + # sharing its cache home; an unspecified cache_path stays global (legacy). + processed_roots: Optional[set[Path]] + if not cache_path: + processed_roots = None + elif _is_processed_dataset_cache_path(repo_id, cache_path): + processed_roots = {Path(cache_path).expanduser().resolve(strict = False).parent} + else: + home = target_root.parent if target_root is not None else None + processed_roots = { + root.resolve(strict = False) + for root in _hf_datasets_cache_roots() + if home is not None and root.resolve(strict = False).parent == home + } + + processed_deleted, processed_failures = _delete_processed_dataset_cache( + repo_id, only_roots = processed_roots + ) failures.extend(processed_failures) if failures: raise HTTPException( @@ -427,15 +448,23 @@ def _delete_cached_dataset_blocking(repo_id: str) -> dict: # ``scan_cache_dir()`` skips blob-only/corrupt repos the revision delete # can't touch, yet the fallback scanner shows them; purge the whole dir. - cache_purged = purge_repo_cache_dirs("dataset", repo_id) - partial_purged = purge_partial_repo("dataset", repo_id) - state_purged = download_manifest.purge_all_state_for_repo("dataset", repo_id) > 0 + # Only for a Hub cache target; a processed-only path has no Hub dir/state. + cache_purged = partial_purged = state_purged = False + if target_root is not None: + cache_purged = purge_repo_cache_dirs("dataset", repo_id, root = target_root) + partial_purged = purge_partial_repo("dataset", repo_id, root = target_root) + state_purged = ( + download_manifest.purge_all_state_for_repo("dataset", repo_id, hub_cache = target_root) + > 0 + ) if not (deleted or processed_deleted or cache_purged or partial_purged or state_purged): raise HTTPException(status_code = 404, detail = "Dataset not found in cache") return {"status": "deleted", "repo_id": repo_id} -def _delete_processed_dataset_cache(repo_id: str) -> tuple[bool, list[str]]: +def _delete_processed_dataset_cache( + repo_id: str, only_roots: Optional[set[Path]] = None +) -> tuple[bool, list[str]]: import shutil target = repo_id.replace("/", "___") @@ -443,6 +472,10 @@ def _delete_processed_dataset_cache(repo_id: str) -> tuple[bool, list[str]]: deleted = False failures: list[str] = [] for root in _hf_datasets_cache_roots(): + # Scope to the selected cache's datasets root(s): a delete must not remove + # processed copies living under other, previously selected cache homes. + if only_roots is not None and root.resolve(strict = False) not in only_roots: + continue try: entries = [ entry diff --git a/studio/backend/hub/services/datasets/downloads.py b/studio/backend/hub/services/datasets/downloads.py index 5efac562fa..b412a339e9 100644 --- a/studio/backend/hub/services/datasets/downloads.py +++ b/studio/backend/hub/services/datasets/downloads.py @@ -159,12 +159,18 @@ async def download_dataset_response( use_xet = download_lifecycle.resolve_effective_use_xet(body.use_xet) transport = download_lifecycle.resolve_transport(use_xet) + from utils.hf_cache_settings import get_hf_cache_paths + + cache_paths = get_hf_cache_paths() + cache_env = cache_paths.child_env({}) claimed, claim_state = _registry.claim( key, transport, repo_type = "dataset", repo_id = repo_id, + hub_cache = str(cache_paths.hub_cache), + xet_cache = str(cache_paths.xet_cache), ) generation = _registry.current_generation(key) if not claimed: @@ -176,7 +182,12 @@ async def download_dataset_response( "accepted": _registry.adoptable(key), "generation": generation, } - download_manifest.clear_cancel_marker("dataset", repo_id, None) + download_manifest.clear_cancel_marker( + "dataset", + repo_id, + None, + hub_cache = cache_paths.hub_cache, + ) state = download_lifecycle.launch_worker( _registry, @@ -185,6 +196,7 @@ async def download_dataset_response( ["--repo-id", repo_id, "--dataset"], hf_token, use_xet = use_xet, + cache_env = cache_env, ), hf_token = hf_token, label = repo_id, diff --git a/studio/backend/hub/services/download_lifecycle.py b/studio/backend/hub/services/download_lifecycle.py index 23f8c7c911..8e14427a56 100644 --- a/studio/backend/hub/services/download_lifecycle.py +++ b/studio/backend/hub/services/download_lifecycle.py @@ -11,7 +11,7 @@ import sys import time import threading from pathlib import Path -from typing import Callable, Optional +from typing import Callable, Mapping, Optional from fastapi import HTTPException @@ -57,6 +57,7 @@ def spawn_worker( *, use_xet: bool, protected_blob_hashes: Optional[frozenset[str]] = None, + cache_env: Optional[Mapping[str, str]] = None, ) -> subprocess.Popen: """Spawn the download worker. @@ -68,7 +69,11 @@ def spawn_worker( """ cwd = backend_dir() mode = download_registry.TRANSPORT_XET if use_xet else download_registry.TRANSPORT_HTTP - env = os.environ.copy() + from utils.hf_cache_settings import get_hf_cache_paths + + env = get_hf_cache_paths().child_env() + if cache_env is not None: + env.update(cache_env) if protected_blob_hashes: env["UNSLOTH_PROTECTED_BLOB_HASHES"] = ",".join(sorted(protected_blob_hashes)) else: @@ -230,6 +235,7 @@ def finalize_worker_exit( (stderr_data or b"").decode("utf-8", "replace").strip(), hf_token = hf_token, ) + metadata = registry.get_job_metadata(key) state = classify_exit(rc, cancel_requested = cancel_requested) if state == "complete": registry.set_job(key, "complete") @@ -252,13 +258,13 @@ def finalize_worker_exit( repo_type, repo_id, download_registry.variant_from_key(key), + hub_cache = metadata.hub_cache if metadata is not None else None, ) except Exception as exc: logger.debug(f"clear_cancel_marker failed for {repo_id} (rc=0): {exc}") elif state == "cancelled": # Read metadata before the terminal set_job so a concurrent eviction # can't drop it; the job key is the fallback variant label. - metadata = registry.get_job_metadata(key) registry.set_job(key, "cancelled") logger.info(f"{log_prefix} cancelled: {label} (rc={rc})") download_registry.persist_cancel_marker( @@ -268,6 +274,7 @@ def finalize_worker_exit( if metadata is not None and metadata.variant else download_registry.variant_from_key(key), cancel_marker_transport or transport, + hub_cache = metadata.hub_cache if metadata is not None else None, logger = logger, ) else: @@ -303,6 +310,7 @@ def _set_retry_failure_state( metadata.transport if metadata is not None and metadata.transport else fallback_transport, + hub_cache = metadata.hub_cache if metadata is not None else None, logger = logger, ) return state @@ -371,6 +379,7 @@ def _try_http_retry( repo_type, repo_id, progress_blob_hashes, + root = Path(original_metadata.hub_cache) if original_metadata.hub_cache else None, ) if progress_blob_hashes else 0 @@ -403,6 +412,8 @@ def _try_http_retry( generation = generation, replace_active = True, cancel_marker_transport = original_metadata.transport, + hub_cache = original_metadata.hub_cache, + xet_cache = original_metadata.xet_cache, ) if claimed: break @@ -446,11 +457,24 @@ def _try_http_retry( label, ) try: + cache_env = ( + { + "HF_HUB_CACHE": original_metadata.hub_cache, + "HF_XET_CACHE": original_metadata.xet_cache, + } + if original_metadata.hub_cache and original_metadata.xet_cache + else None + ) + spawn_kwargs = { + "use_xet": False, + "protected_blob_hashes": peer_hashes or None, + } + if cache_env is not None: + spawn_kwargs["cache_env"] = cache_env proc = spawn_worker( args, hf_token, - use_xet = False, - protected_blob_hashes = peer_hashes or None, + **spawn_kwargs, ) except Exception as exc: scrubbed = download_registry.scrub_secrets(str(exc), hf_token = hf_token) diff --git a/studio/backend/hub/services/models/cache_inventory.py b/studio/backend/hub/services/models/cache_inventory.py index c1b864bb63..76dd2337aa 100644 --- a/studio/backend/hub/services/models/cache_inventory.py +++ b/studio/backend/hub/services/models/cache_inventory.py @@ -34,7 +34,6 @@ from hub.services.models.common import ( _is_mmproj_filename, _is_transformers_safetensors_weight_name, _local_inventory_id, - _prefer_complete_larger, _runtime_for_format, ) @@ -250,24 +249,46 @@ def _repo_gguf_blob_map(repo_info, *, include_companions: bool = False) -> dict[ def _prefer_cache_row(candidate: dict, existing: Optional[dict]) -> bool: if existing is None: return True - return _prefer_complete_larger( - bool(candidate.get("partial")), - int(candidate.get("size_bytes") or 0), - bool(existing.get("partial")), - int(existing.get("size_bytes") or 0), - ) + candidate_partial = bool(candidate.get("partial")) + existing_partial = bool(existing.get("partial")) + if candidate_partial != existing_partial: + return not candidate_partial + candidate_active = bool(candidate.get("active_cache")) + existing_active = bool(existing.get("active_cache")) + if candidate_active != existing_active: + return candidate_active + return int(candidate.get("size_bytes") or 0) > int(existing.get("size_bytes") or 0) def _cache_inventory_fields( repo_id: str, model_format: ModelFormat, *, + repo_path: Optional[Path] = None, + snapshot_path: Optional[Path] = None, + active_hub_cache: Optional[Path] = None, partial: bool = False, requires_variant: bool = False, ) -> dict: + load_id = repo_id + active_cache = True + if repo_path is not None: + try: + if active_hub_cache is None: + from utils.hf_cache_settings import get_hf_cache_paths + active_hub_cache = get_hf_cache_paths().hub_cache + active_root = active_hub_cache.resolve(strict = False) + cached_root = repo_path.parent.resolve(strict = False) + if cached_root != active_root: + active_cache = False + load_id = str(snapshot_path or repo_path.resolve(strict = False)) + except (OSError, RuntimeError, ValueError): + active_cache = False + load_id = str(snapshot_path or repo_path) return { "inventory_id": _local_inventory_id("cache", model_format, repo_id), - "load_id": repo_id, + "load_id": load_id, + "active_cache": active_cache, "model_format": model_format, "runtime": _runtime_for_format(model_format), "format_variant": None, @@ -294,6 +315,9 @@ def _is_hidden_infra_repo(*values: str | None) -> bool: def _scan_cached_gguf() -> list[dict]: """Synchronous HF-cache disk walk for GGUF repos; runs in a worker thread.""" cache_scans = all_hf_cache_scans() + from utils.hf_cache_settings import get_hf_cache_paths + + active_hub_cache = get_hf_cache_paths().hub_cache seen_lower: dict[str, dict] = {} for hf_cache in cache_scans: @@ -305,7 +329,10 @@ def _scan_cached_gguf() -> list[dict]: repo_path = Path(repo_info.repo_path) snapshot_path = _cached_model_snapshot_path(repo_path) total_size = _repo_gguf_size_bytes(repo_info) - has_variant_state, variant_state_size = _gguf_variant_state_summary(repo_id) + has_variant_state, variant_state_size = _gguf_variant_state_summary( + repo_id, + hub_cache = repo_path.parent, + ) is_hidden_infra = _is_hidden_infra_repo( repo_id, str(repo_path), @@ -342,6 +369,9 @@ def _scan_cached_gguf() -> list[dict]: _cache_inventory_fields( repo_id, "gguf", + repo_path = repo_path, + snapshot_path = snapshot_path, + active_hub_cache = active_hub_cache, partial = bool(row["partial"]), requires_variant = True, ) @@ -543,6 +573,9 @@ def _cached_model_local_metadata(repo_path: Path) -> dict: def _scan_cached_models() -> list[dict]: """Synchronous HF-cache disk walk for non-GGUF model repos; runs in a worker thread.""" cache_scans = all_hf_cache_scans() + from utils.hf_cache_settings import get_hf_cache_paths + + active_hub_cache = get_hf_cache_paths().hub_cache seen_lower: dict[str, dict] = {} inspected = 0 @@ -606,6 +639,9 @@ def _scan_cached_models() -> list[dict]: _cache_inventory_fields( repo_id, payload.model_format, + repo_path = repo_path, + snapshot_path = snapshot_path, + active_hub_cache = active_hub_cache, partial = bool(row["partial"]), ) ) diff --git a/studio/backend/hub/services/models/common.py b/studio/backend/hub/services/models/common.py index f381bffe9c..4c0e296fdc 100644 --- a/studio/backend/hub/services/models/common.py +++ b/studio/backend/hub/services/models/common.py @@ -150,7 +150,9 @@ def _prefer_complete_larger( return candidate_size_bytes > existing_size_bytes -def _gguf_variant_state_summary(repo_id: str) -> tuple[bool, int]: +def _gguf_variant_state_summary( + repo_id: str, *, hub_cache: Optional[str | Path] = None +) -> tuple[bool, int]: """Whether GGUF variant-scoped state exists and its expected size; a cancelled/in-progress variant may have only manifests/markers/`.incomplete` blobs, which inventory needs to avoid a generic fallback row.""" from hub.utils import download_manifest @@ -159,10 +161,16 @@ def _gguf_variant_state_summary(repo_id: str) -> tuple[bool, int]: for variant, _path in download_manifest.iter_variant_manifests( "model", repo_id, + hub_cache = hub_cache, ): key = variant.lower() variant_keys.add(key) - manifest = download_manifest.read_manifest("model", repo_id, variant) + manifest = download_manifest.read_manifest( + "model", + repo_id, + variant, + hub_cache = hub_cache, + ) if manifest is None: continue size_by_variant[key] = max( @@ -172,6 +180,7 @@ def _gguf_variant_state_summary(repo_id: str) -> tuple[bool, int]: for variant, _path in download_manifest.iter_variant_markers( "model", repo_id, + hub_cache = hub_cache, ): variant_keys.add(variant.lower()) return bool(variant_keys), sum(size_by_variant.values()) @@ -432,8 +441,13 @@ def _local_model_info( base_model_source: Optional[str] = None, adapter_type: Optional[str] = None, training_method: Optional[str] = None, + active_cache: Optional[bool] = None, ) -> LocalModelInfo: - load_id = model_id if source == "hf_cache" and model_id else str(load_path) + load_id = ( + model_id + if source == "hf_cache" and model_id and active_cache is not False + else str(load_path) + ) semantic_id = model_id or str(load_path) return LocalModelInfo( id = load_id, @@ -445,6 +459,7 @@ def _local_model_info( ), load_id = load_id, model_id = model_id, + active_cache = active_cache if source == "hf_cache" else None, display_name = display_name or (scan_path.stem if scan_path.is_file() else scan_path.name), path = str(load_path), size_bytes = max(0, int(size_bytes or 0)), @@ -476,6 +491,7 @@ def _classify_local_path( model_id: Optional[str] = None, updated_at: Optional[float] = None, partial: bool = False, + active_cache: Optional[bool] = None, ) -> list[LocalModelInfo]: load_path = load_path or scan_path files = ( @@ -512,6 +528,7 @@ def _classify_local_path( requires_variant = scan_path.is_dir(), format_variant = variant, size_bytes = gguf_size_bytes, + active_cache = active_cache, ) ) @@ -574,6 +591,7 @@ def _classify_local_path( ), adapter_type = adapter_type if model_format == "adapter" else None, training_method = training_method if model_format == "adapter" else None, + active_cache = active_cache, ) ) elif not rows: @@ -592,6 +610,7 @@ def _classify_local_path( updated_at = updated_at, partial = partial or trusted_hf_cache_repo, size_bytes = size_bytes, + active_cache = active_cache, ) ) diff --git a/studio/backend/hub/services/models/deletion.py b/studio/backend/hub/services/models/deletion.py index 636a223d4e..c736908058 100644 --- a/studio/backend/hub/services/models/deletion.py +++ b/studio/backend/hub/services/models/deletion.py @@ -19,8 +19,10 @@ from hub.utils import inventory_scan as hf_cache_scan from hub.utils.gguf import extract_quant_label, extract_quant_token from hub.utils.hf_cache_state import ( INCOMPLETE_SUFFIX, + iter_repo_cache_dirs, purge_partial_repo, purge_repo_cache_dirs, + resolve_delete_target_root, ) from hub.utils.paths import ( is_valid_gguf_variant as _is_valid_gguf_variant, @@ -184,6 +186,7 @@ def _delete_gguf_variant_from_repos( hf_token: Optional[str], *, sibling_active: bool = False, + root: Optional[Path] = None, ) -> dict: failures: list[str] = [] removed_snapshots = 0 @@ -265,6 +268,7 @@ def _delete_gguf_variant_from_repos( hf_token, extra_hashes = frozenset(completed_hashes), companions = not sibling_active, + root = root, ) if incomplete_result.unresolved: raise HTTPException( @@ -276,7 +280,7 @@ def _delete_gguf_variant_from_repos( ), ) - state_purged = download_manifest.purge_state("model", repo_id, variant) + state_purged = download_manifest.purge_state("model", repo_id, variant, hub_cache = root) # Reclaim the empty quant folder so it stops 404ing on delete. removed_dirs, dir_failures = _remove_empty_variant_dirs(target_repos, variant) removed_snap_dirs, snap_dir_failures = _remove_empty_snapshot_dirs(target_repos) @@ -316,6 +320,8 @@ def reclaim_replaced_gguf_variant( variant: str, keep_main_hashes: frozenset[str], hf_token: Optional[str] = None, + *, + hub_cache: Optional[str | Path] = None, ) -> dict: """Prune stale main-GGUF files for a variant after a replacement verified. @@ -366,12 +372,22 @@ def reclaim_replaced_gguf_variant( "reason": "scan_failed", } + if hub_cache is None: + from utils.hf_cache_settings import get_hf_cache_paths + hub_cache = get_hf_cache_paths().hub_cache + try: + target_hub_cache = Path(hub_cache).expanduser().resolve(strict = False) + except (OSError, RuntimeError, ValueError): + target_hub_cache = Path(hub_cache).expanduser() + candidate_repos = [ repo_info for hf_cache in cache_scans for repo_info in hf_cache.repos if str(getattr(repo_info, "repo_type", "")) == "model" and str(getattr(repo_info, "repo_id", "")).lower() == repo_id.lower() + and getattr(repo_info, "repo_path", None) + and Path(repo_info.repo_path).parent.resolve(strict = False) == target_hub_cache ] try: matched_repo_ids = resolve_destructive_repo_ids( @@ -493,10 +509,24 @@ def reclaim_replaced_gguf_variant( def _loaded_id_matches_repo(loaded_id: str, repo_id: str) -> bool: - """True when *loaded_id* is *repo_id* or a file within it; ``/``-boundary aware so ``org/model`` doesn't match sibling ``org/model-v2``.""" + """Match a loaded repo ID or an on-disk path inside any copy of the repo.""" rid = repo_id.lower() lid = loaded_id.lower() - return lid == rid or lid.startswith(f"{rid}/") + if lid == rid or lid.startswith(f"{rid}/"): + return True + + try: + loaded_path = Path(loaded_id).expanduser().resolve(strict = False) + except (OSError, RuntimeError, ValueError): + return False + for repo_dir in iter_repo_cache_dirs("model", repo_id): + try: + resolved_repo = repo_dir.resolve(strict = False) + if loaded_path == resolved_repo or loaded_path.is_relative_to(resolved_repo): + return True + except (OSError, RuntimeError, ValueError): + continue + return False def _loaded_repo_variant_blocks_delete( @@ -560,6 +590,7 @@ async def delete_cached_model_response( repo_id: str, variant: Optional[str] = None, hf_token: Optional[str] = None, + cache_path: Optional[str] = None, ): """Delete a cached model repo (or a specific GGUF variant) from the HF cache. @@ -603,14 +634,19 @@ async def delete_cached_model_response( ) raise HTTPException(status_code = 400, detail = detail) try: - return await asyncio.to_thread(_delete_cached_model_blocking, repo_id, variant, hf_token) + return await asyncio.to_thread( + _delete_cached_model_blocking, repo_id, variant, hf_token, cache_path + ) finally: downloads.registry.end_delete(repo_key, variant) cache_inventory.invalidate_hf_cache_scans() def _delete_cached_model_blocking( - repo_id: str, variant: Optional[str], hf_token: Optional[str] + repo_id: str, + variant: Optional[str], + hf_token: Optional[str], + cache_path: Optional[str] = None, ) -> dict: try: # If a sibling quant is downloading concurrently, restrict this delete to @@ -621,13 +657,26 @@ def _delete_cached_model_blocking( cache_scans = cache_inventory.all_hf_cache_scans() - candidate_entries = [] + # A repo can live in several remembered caches. Group its copies by the + # cache root that owns each, then target exactly one cache so a delete + # never removes copies in other, previously selected caches. + owners: dict = {} for hf_cache in cache_scans: for repo_info in hf_cache.repos: if str(repo_info.repo_type) != "model": continue - if repo_info.repo_id.lower() == repo_id.lower(): - candidate_entries.append((hf_cache, repo_info)) + if repo_info.repo_id.lower() != repo_id.lower(): + continue + try: + owner = Path(repo_info.repo_path).parent.resolve(strict = False) + except (OSError, RuntimeError, ValueError): + continue + owners.setdefault(owner, []).append((hf_cache, repo_info)) + + target_root = resolve_delete_target_root("model", repo_id, cache_path, owners.keys()) + if target_root is None: + raise HTTPException(status_code = 400, detail = "Invalid cache_path") + candidate_entries = owners.get(target_root, []) matched_repo_ids = resolve_destructive_repo_ids( repo_id, @@ -642,10 +691,15 @@ def _delete_cached_model_blocking( if not target_entries: if variant is None: - cache_purged = purge_repo_cache_dirs("model", repo_id) or purge_partial_repo( - "model", repo_id + cache_purged = purge_repo_cache_dirs( + "model", repo_id, root = target_root + ) or purge_partial_repo("model", repo_id, root = target_root) + state_purged = ( + download_manifest.purge_all_state_for_repo( + "model", repo_id, hub_cache = target_root + ) + > 0 ) - state_purged = download_manifest.purge_all_state_for_repo("model", repo_id) > 0 if cache_purged or state_purged: return {"status": "deleted", "repo_id": repo_id} if variant: @@ -654,6 +708,7 @@ def _delete_cached_model_blocking( variant, hf_token, companions = not sibling_active, + root = target_root, ) if incomplete_result.unresolved: raise HTTPException( @@ -668,6 +723,7 @@ def _delete_cached_model_blocking( "model", repo_id, variant, + hub_cache = target_root, ) if incomplete_result.deleted > 0 or state_purged: return { @@ -684,6 +740,7 @@ def _delete_cached_model_blocking( [repo for _cache, repo in target_entries], hf_token, sibling_active = sibling_active, + root = target_root, ) deleted_revisions = False @@ -702,9 +759,11 @@ def _delete_cached_model_blocking( delete_strategy.execute() deleted_revisions = True - cache_purged = purge_repo_cache_dirs("model", repo_id) - partial_purged = purge_partial_repo("model", repo_id) - state_purged = download_manifest.purge_all_state_for_repo("model", repo_id) > 0 + cache_purged = purge_repo_cache_dirs("model", repo_id, root = target_root) + partial_purged = purge_partial_repo("model", repo_id, root = target_root) + state_purged = ( + download_manifest.purge_all_state_for_repo("model", repo_id, hub_cache = target_root) > 0 + ) if not (deleted_revisions or cache_purged or partial_purged or state_purged): raise HTTPException(status_code = 404, detail = "No revisions found for model") diff --git a/studio/backend/hub/services/models/downloads.py b/studio/backend/hub/services/models/downloads.py index 862af0141a..c93b21c082 100644 --- a/studio/backend/hub/services/models/downloads.py +++ b/studio/backend/hub/services/models/downloads.py @@ -90,6 +90,7 @@ def _spawn_download_worker( hf_token: Optional[str], use_xet: bool = True, protected_blob_hashes: Optional[frozenset[str]] = None, + cache_env: Optional[dict[str, str]] = None, ) -> subprocess.Popen: args = ["--repo-id", repo_id] if variant: @@ -99,6 +100,7 @@ def _spawn_download_worker( hf_token, use_xet = use_xet, protected_blob_hashes = protected_blob_hashes, + cache_env = cache_env, ) @@ -125,6 +127,10 @@ async def download_model_response(body: DownloadModelRequest, hf_token: Optional key = _download_job_key(repo_id, variant) use_xet = download_lifecycle.resolve_effective_use_xet(body.use_xet) transport = download_lifecycle.resolve_transport(use_xet) + from utils.hf_cache_settings import get_hf_cache_paths + + cache_paths = get_hf_cache_paths() + cache_env = cache_paths.child_env({}) variant_blob_hashes = frozenset() variant_progress_blob_hashes = frozenset() completed_baseline_bytes = 0 @@ -175,6 +181,8 @@ async def download_model_response(body: DownloadModelRequest, hf_token: Optional progress_blob_hashes = variant_progress_blob_hashes, completed_baseline_bytes = completed_baseline_bytes, admission_check = lambda: not _load_in_flight(repo_id), + hub_cache = str(cache_paths.hub_cache), + xet_cache = str(cache_paths.xet_cache), ) generation = _registry.current_generation(key) if not claimed: @@ -189,7 +197,12 @@ async def download_model_response(body: DownloadModelRequest, hf_token: Optional "accepted": _registry.adoptable(key), "generation": generation, } - download_manifest.clear_cancel_marker("model", repo_id, variant) + download_manifest.clear_cancel_marker( + "model", + repo_id, + variant, + hub_cache = cache_paths.hub_cache, + ) # Blobs a concurrent same-repo variant is already writing (e.g. a shared # mmproj). The worker must not purge these during cache preparation. protected_blob_hashes = _registry.peer_blob_hashes(key) if variant else frozenset() @@ -204,6 +217,7 @@ async def download_model_response(body: DownloadModelRequest, hf_token: Optional hf_token, use_xet = use_xet, protected_blob_hashes = protected_blob_hashes, + cache_env = cache_env, ), hf_token = hf_token, label = label, diff --git a/studio/backend/hub/services/models/folder_browser.py b/studio/backend/hub/services/models/folder_browser.py index 7d9c3ac665..effb6a32ae 100644 --- a/studio/backend/hub/services/models/folder_browser.py +++ b/studio/backend/hub/services/models/folder_browser.py @@ -30,6 +30,7 @@ from hub.utils.paths import ( ) from utils.paths.external_media import ( linux_run_media_mount_roots, + macos_volume_roots, windows_drive_roots, ) from hub.services.models.common import _safe_is_dir @@ -187,7 +188,7 @@ def _build_browse_allowlist( _add(Path.home()) if media_roots is None: - media_roots = linux_run_media_mount_roots() + media_roots = [*linux_run_media_mount_roots(), *macos_volume_roots()] if drive_roots is None: drive_roots = windows_drive_roots() for p in media_roots: @@ -195,6 +196,12 @@ def _build_browse_allowlist( for p in drive_roots: _add(p) _add(_resolve_hf_cache_dir()) + try: + from utils.hf_cache_settings import known_hf_cache_homes + for cache_home in known_hf_cache_homes(): + _add(cache_home) + except Exception: # noqa: BLE001 -- best-effort + pass try: _add(hf_default_cache_dir()) except Exception: # noqa: BLE001 -- best-effort @@ -431,7 +438,7 @@ def browse_folders_response( # Probe removable-media and Windows drive roots once; the allowlist and # chips reuse the result so a disconnected mapped drive isn't scanned twice. - media_roots = linux_run_media_mount_roots() + media_roots = [*linux_run_media_mount_roots(), *macos_volume_roots()] drive_roots = windows_drive_roots() # Build the allowlist once -- the sandbox check and suggestion chips share # it so chips are always navigable. diff --git a/studio/backend/hub/services/models/gguf_variants.py b/studio/backend/hub/services/models/gguf_variants.py index 33f0297ff5..533fd2ca5a 100644 --- a/studio/backend/hub/services/models/gguf_variants.py +++ b/studio/backend/hub/services/models/gguf_variants.py @@ -9,6 +9,7 @@ import asyncio import threading import time from collections import OrderedDict +from pathlib import Path from typing import NamedTuple, Optional from fastapi import HTTPException @@ -22,6 +23,7 @@ from hub.utils.hf_errors import hf_error_status from hub.utils.hf_cache_state import ( INCOMPLETE_SUFFIX, iter_destructive_repo_cache_dirs, + repo_cache_dir_name, ) from hub.utils.gguf import ( extract_quant_label, @@ -233,8 +235,14 @@ def _manifest_variant_blob_hashes( variant: str, *, include_companions: bool = True, + repo_cache_dir: Optional[Path] = None, ) -> frozenset[str]: - manifest = download_manifest.read_manifest("model", repo_id, variant) + manifest = download_manifest.read_manifest( + "model", + repo_id, + variant, + hub_cache = repo_cache_dir.parent if repo_cache_dir is not None else None, + ) if manifest is None: return frozenset() variant_key = variant.lower() @@ -257,6 +265,7 @@ def gguf_variant_blob_hashes( *, include_companions: bool = True, allow_remote: bool = True, + repo_cache_dir: Optional[Path] = None, ) -> frozenset[str]: key = _variant_blob_hash_cache_key( repo_id, @@ -271,9 +280,9 @@ def gguf_variant_blob_hashes( repo_id, variant, include_companions = include_companions, + repo_cache_dir = repo_cache_dir, ) if hashes: - _variant_hash_cache_set(key, hashes) return hashes requirement_key = _variant_hash_cache_key(repo_id, variant, hf_token) requirement = _variant_requirement_cache_get(requirement_key) @@ -287,11 +296,22 @@ def gguf_variant_blob_hashes( return frozenset() -def _partial_transport_for_variant(repo_id: str, variant: str) -> Optional[str]: - return hf_cache_scan.partial_transport_for("model", repo_id, variant) +def _partial_transport_for_variant( + repo_id: str, + variant: str, + repo_cache_dir: Optional[Path] = None, +) -> Optional[str]: + return hf_cache_scan.partial_transport_for( + "model", + repo_id, + variant, + repo_cache_dir, + ) -def _local_main_gguf_blobs_by_quant(repo_id: str) -> dict[str, dict[str, set[str]]]: +def _local_main_gguf_blobs_by_quant( + repo_id: str, repo_cache_dir: Optional[Path] = None +) -> dict[str, dict[str, set[str]]]: """Map quant -> repo-relative expected GGUF filename -> cached blob hashes. Shared companions are copied into each main-quant bucket so update checks can @@ -313,6 +333,14 @@ def _local_main_gguf_blobs_by_quant(repo_id: str) -> dict[str, dict[str, set[str continue if str(getattr(repo_info, "repo_id", "")).lower() != target_lower: continue + if repo_cache_dir is not None: + try: + if Path(repo_info.repo_path).resolve(strict = False) != repo_cache_dir.resolve( + strict = False + ): + continue + except (AttributeError, OSError, RuntimeError, ValueError): + continue for path, hashes in cache_inventory._repo_gguf_blob_map( repo_info, include_companions = True, @@ -388,6 +416,7 @@ def delete_variant_incomplete_blobs_result( *, extra_hashes: frozenset[str] = frozenset(), companions: bool = True, + root: Optional[Path] = None, ) -> VariantIncompleteDeleteResult: # With a sibling still downloading, ``companions=False`` keeps a shared mmproj # from being unlinked out from under it; the repo's last delete reclaims it. @@ -409,8 +438,9 @@ def delete_variant_incomplete_blobs_result( ) deleted = 0 # Destructive iterator: only the exact-case match (or abort if ambiguous), - # so a case-variant sibling repo's partials are never unlinked. - for entry in iter_destructive_repo_cache_dirs("model", repo_id): + # so a case-variant sibling repo's partials are never unlinked. ``root`` scopes + # the purge to one cache so a delete never touches another cache's partials. + for entry in iter_destructive_repo_cache_dirs("model", repo_id, root = root): blobs_dir = entry / "blobs" if not blobs_dir.is_dir(): continue @@ -425,15 +455,37 @@ def delete_variant_incomplete_blobs_result( return VariantIncompleteDeleteResult(deleted = deleted, unresolved = False) +def _repo_cache_dir_for_request(repo_id: str, local_path: Optional[str]) -> Path: + """Resolve the one Hub repo cache represented by this variant request.""" + expected_name = repo_cache_dir_name("model", repo_id).lower() + if local_path: + try: + local = Path(local_path).expanduser().resolve(strict = False) + for candidate in (local, *local.parents): + if candidate.name.lower() == expected_name: + return candidate + except (OSError, RuntimeError, ValueError): + pass + from utils.hf_cache_settings import get_hf_cache_paths + + return get_hf_cache_paths().hub_cache / repo_cache_dir_name("model", repo_id) + + def _mark_empty_dir_cleanables( - repo_id: str, response: GgufVariantsResponse + repo_id: str, + response: GgufVariantsResponse, + repo_cache_dir: Optional[Path] = None, ) -> GgufVariantsResponse: """Surface empty leftover ``/`` folders (interrupted downloads) as partial so the UI can delete them -- on local/offline paths too, not just a remote listing. A listed quant is flipped to partial; an unlisted one is appended as a zero-byte cleanable entry.""" try: - empty_labels = list_empty_gguf_variant_dirs(repo_id) + empty_labels = ( + list_empty_gguf_variant_dirs(repo_id, root = repo_cache_dir.parent) + if repo_cache_dir is not None + else list_empty_gguf_variant_dirs(repo_id) + ) except Exception as e: logger.warning(f"Failed to scan empty GGUF variant folders for {repo_id}: {e}") return response @@ -468,6 +520,11 @@ async def get_gguf_variants_response( """ def _compute() -> GgufVariantsResponse: + repo_cache_dir = ( + None if is_local_path(repo_id) else _repo_cache_dir_for_request(repo_id, local_path) + ) + hub_cache = repo_cache_dir.parent if repo_cache_dir is not None else None + def _local_response( response_repo_id: str, variants, has_vision: bool ) -> GgufVariantsResponse: @@ -511,6 +568,7 @@ async def get_gguf_variants_response( partial_transport = _partial_transport_for_variant( response_repo_id, v.quant, + repo_cache_dir, ), ) for v in variants @@ -532,7 +590,7 @@ async def get_gguf_variants_response( local_only = prefer_local_cache or offline if local_only: - cached = list_gguf_variants_from_hf_cache(repo_id) + cached = list_gguf_variants_from_hf_cache(repo_id, root = hub_cache) if cached is not None: variants, has_vision = cached return _local_response(repo_id, variants, has_vision) @@ -540,7 +598,7 @@ async def get_gguf_variants_response( variants, has_vision = list_local_gguf_variants(local_path) if variants or has_vision: return _local_response(repo_id, variants, has_vision) - partial = list_partial_gguf_variants_from_state(repo_id) + partial = list_partial_gguf_variants_from_state(repo_id, hub_cache = hub_cache) if partial is not None: variants, has_vision = partial return _partial_local_response(repo_id, variants, has_vision) @@ -560,11 +618,11 @@ async def get_gguf_variants_response( try: variants, has_vision, siblings = list_gguf_variants(repo_id, hf_token = hf_token) except Exception: - cached = list_gguf_variants_from_hf_cache(repo_id) + cached = list_gguf_variants_from_hf_cache(repo_id, root = hub_cache) if cached is not None: variants, has_vision = cached return _local_response(repo_id, variants, has_vision) - partial = list_partial_gguf_variants_from_state(repo_id) + partial = list_partial_gguf_variants_from_state(repo_id, hub_cache = hub_cache) if partial is not None: variants, has_vision = partial return _partial_local_response(repo_id, variants, has_vision) @@ -581,7 +639,7 @@ async def get_gguf_variants_response( cached_filenames_by_snapshot: list[dict[str, int]] = [] cached_quant_bytes_by_snapshot: list[dict[str, int]] = [] if _is_valid_repo_id(repo_id): - for snap in iter_hf_cache_snapshots(repo_id): + for snap in iter_hf_cache_snapshots(repo_id, root = hub_cache): try: gguf_paths = list(_iter_gguf_paths(snap)) except (OSError, RuntimeError, ValueError) as e: @@ -694,11 +752,20 @@ async def get_gguf_variants_response( partial_quants: set[str] = set() partial_quant_transports: dict[str, Optional[str]] = {} try: - incomplete_hashes = download_registry.incomplete_blob_hashes("model", repo_id) + incomplete_hashes = download_registry.incomplete_blob_hashes( + "model", + repo_id, + active_only = True, + root = hub_cache, + ) except Exception as e: logger.warning(f"Failed to compute partial GGUF variants for {repo_id}: {e}") incomplete_hashes = set() - scan_snapshot_dir = hf_cache_scan.resolve_snapshot_dir_for_scan("model", repo_id) + scan_snapshot_dir = hf_cache_scan.resolve_snapshot_dir_for_scan( + "model", + repo_id, + repo_cache_dir, + ) # Manifest + marker + main incomplete-blob check: catches variants whose # download was cancelled or whose expected shards are missing/undersized. for variant in variants: @@ -711,6 +778,7 @@ async def get_gguf_variants_response( variant.quant, hf_token, include_companions = False, + repo_cache_dir = repo_cache_dir, ) if hf_cache_scan.is_variant_partial( repo_id, @@ -718,11 +786,13 @@ async def get_gguf_variants_response( scan_snapshot_dir, incomplete_blob_hashes = incomplete_hashes, variant_blob_hashes = variant_hashes, + repo_cache_dir = repo_cache_dir, ): partial_quants.add(variant.quant) partial_quant_transports[variant.quant] = _partial_transport_for_variant( repo_id, variant.quant, + repo_cache_dir, ) except Exception as e: logger.warning( @@ -744,10 +814,14 @@ async def get_gguf_variants_response( partial_quants.add(variant.quant) partial_quant_transports.setdefault( variant.quant, - _partial_transport_for_variant(repo_id, variant.quant), + _partial_transport_for_variant( + repo_id, + variant.quant, + repo_cache_dir, + ), ) - local_blobs_by_quant = _local_main_gguf_blobs_by_quant(repo_id) + local_blobs_by_quant = _local_main_gguf_blobs_by_quant(repo_id, repo_cache_dir) def _variant_detail(v) -> GgufVariantDetail: is_partial = v.quant in partial_quants @@ -790,14 +864,20 @@ async def get_gguf_variants_response( if skip: raise enriched = _mark_empty_dir_cleanables( - repo_id, GgufVariantsResponse(repo_id = repo_id, variants = []) + repo_id, + GgufVariantsResponse(repo_id = repo_id, variants = []), + _repo_cache_dir_for_request(repo_id, local_path), ) if enriched.variants: return enriched raise if skip: return response - return _mark_empty_dir_cleanables(repo_id, response) + return _mark_empty_dir_cleanables( + repo_id, + response, + _repo_cache_dir_for_request(repo_id, local_path), + ) try: return await asyncio.to_thread(_compute_with_cleanables) diff --git a/studio/backend/hub/services/models/local_inventory.py b/studio/backend/hub/services/models/local_inventory.py index b34532fa35..9cf260b157 100644 --- a/studio/backend/hub/services/models/local_inventory.py +++ b/studio/backend/hub/services/models/local_inventory.py @@ -106,11 +106,8 @@ def _is_model_directory_for_scan(path: Path, *, entry_limit: int | None) -> bool def _resolve_hf_cache_dir() -> Path: - try: - from huggingface_hub.constants import HF_HUB_CACHE - return Path(HF_HUB_CACHE) - except Exception: - return Path.home() / ".cache" / "huggingface" / "hub" + from utils.hf_cache_settings import get_hf_cache_paths + return get_hf_cache_paths().hub_cache def _scan_models_dir( @@ -202,7 +199,12 @@ def _hf_repo_dir_has_content(repo_dir: Path) -> bool: return False -def _scan_hf_cache(cache_dir: Path, *, entry_limit: int | None = None) -> List[LocalModelInfo]: +def _scan_hf_cache( + cache_dir: Path, + *, + entry_limit: int | None = None, + active_cache: bool = True, +) -> List[LocalModelInfo]: if not _safe_is_dir(cache_dir): return [] @@ -240,7 +242,10 @@ def _scan_hf_cache(cache_dir: Path, *, entry_limit: int | None = None) -> List[L repo_dir, ) gguf_partial = hf_cache_scan.is_gguf_repo_partial(model_id, repo_dir) - has_gguf_variant_state, gguf_variant_state_size = _gguf_variant_state_summary(model_id) + has_gguf_variant_state, gguf_variant_state_size = _gguf_variant_state_summary( + model_id, + hub_cache = cache_dir, + ) snapshot_partial_transport = ( hf_cache_scan.partial_transport_for( "model", @@ -252,23 +257,25 @@ def _scan_hf_cache(cache_dir: Path, *, entry_limit: int | None = None) -> List[L ) resolved = hf_cache_scan.resolve_hf_cache_realpath(repo_dir) scan_path = Path(resolved) if resolved else repo_dir + load_path = repo_dir if active_cache else scan_path # partial=False here; _apply_format_aware_partial below rewrites per-row # so a hybrid repo's gguf row doesn't taint its safetensors row. rows = _classify_local_path( scan_path, "hf_cache", - load_path = repo_dir, + load_path = load_path, display_name = model_id.split("/")[-1], model_id = model_id, updated_at = updated_at, partial = False, + active_cache = active_cache, ) if not rows: if has_gguf_variant_state and gguf_partial: rows = [ _local_model_info( scan_path = repo_dir, - load_path = repo_dir, + load_path = load_path, source = "hf_cache", model_format = "gguf", display_name = model_id.split("/")[-1], @@ -277,6 +284,7 @@ def _scan_hf_cache(cache_dir: Path, *, entry_limit: int | None = None) -> List[L partial = True, requires_variant = True, size_bytes = gguf_variant_state_size, + active_cache = active_cache, ) ] else: @@ -285,13 +293,14 @@ def _scan_hf_cache(cache_dir: Path, *, entry_limit: int | None = None) -> List[L rows = [ _local_model_info( scan_path = repo_dir, - load_path = repo_dir, + load_path = load_path, source = "hf_cache", model_format = "unknown", display_name = model_id.split("/")[-1], model_id = model_id, updated_at = updated_at, partial = snapshot_partial or gguf_partial, + active_cache = active_cache, ) ] elif ( @@ -302,7 +311,7 @@ def _scan_hf_cache(cache_dir: Path, *, entry_limit: int | None = None) -> List[L rows.append( _local_model_info( scan_path = repo_dir, - load_path = repo_dir, + load_path = load_path, source = "hf_cache", model_format = "gguf", display_name = model_id.split("/")[-1], @@ -311,6 +320,7 @@ def _scan_hf_cache(cache_dir: Path, *, entry_limit: int | None = None) -> List[L partial = True, requires_variant = True, size_bytes = gguf_variant_state_size, + active_cache = active_cache, ) ) rows = _apply_format_aware_partial( @@ -515,14 +525,39 @@ async def _collect_models_from_default_sources( local_models += await _scan_source("HF cache", _scan_hf_cache, hf_cache_dir) if _safe_is_dir(legacy_hf) and legacy_hf.resolve() != hf_cache_dir.resolve(): - local_models += await _scan_source("legacy HF cache", _scan_hf_cache, legacy_hf) + local_models += await _scan_source( + "legacy HF cache", + lambda path: _scan_hf_cache(path, active_cache = False), + legacy_hf, + ) if ( _safe_is_dir(hf_default) and hf_default.resolve() != hf_cache_dir.resolve() and hf_default.resolve() != legacy_hf.resolve() ): - local_models += await _scan_source("default HF cache", _scan_hf_cache, hf_default) + local_models += await _scan_source( + "default HF cache", + lambda path: _scan_hf_cache(path, active_cache = False), + hf_default, + ) + + from utils.hf_cache_settings import known_hf_hub_caches + + seen_hf = { + os.path.normcase(str(path.resolve(strict = False))) + for path in (hf_cache_dir, legacy_hf, hf_default) + } + for previous_cache in known_hf_hub_caches(): + key = os.path.normcase(str(previous_cache.resolve(strict = False))) + if key in seen_hf: + continue + seen_hf.add(key) + local_models += await _scan_source( + "previous HF cache", + lambda path: _scan_hf_cache(path, active_cache = False), + previous_cache, + ) for lm_dir in lm_dirs: local_models += await _scan_source("LM Studio", _scan_lmstudio_dir, lm_dir) @@ -543,7 +578,11 @@ def _scan_custom_folder(folder_path: Path) -> List[LocalModelInfo]: limit = _MAX_MODELS_PER_CUSTOM_FOLDER, entry_limit = _MAX_CUSTOM_FOLDER_ENTRIES, ) - + _scan_hf_cache(folder_path, entry_limit = _MAX_CUSTOM_FOLDER_ENTRIES) + + _scan_hf_cache( + folder_path, + entry_limit = _MAX_CUSTOM_FOLDER_ENTRIES, + active_cache = False, + ) + _scan_lmstudio_dir(folder_path, entry_limit = _MAX_CUSTOM_FOLDER_ENTRIES) ) if m.model_format in supported_formats @@ -610,12 +649,20 @@ def _dedupe_local_models(local_models: List[LocalModelInfo]) -> list[LocalModelI row_key = model.inventory_id or model.id key = f"{row_key}\x00custom" if model.source == "custom" else row_key existing = deduped.get(key) - if existing is None or _prefer_complete_larger( - model.partial, - model.size_bytes, - existing.partial, - existing.size_bytes, - ): + prefer_candidate = existing is None + if existing is not None: + if model.partial != existing.partial: + prefer_candidate = not model.partial + elif (model.active_cache is True) != (existing.active_cache is True): + prefer_candidate = model.active_cache is True + else: + prefer_candidate = _prefer_complete_larger( + model.partial, + model.size_bytes, + existing.partial, + existing.size_bytes, + ) + if prefer_candidate: deduped[key] = model return sorted( deduped.values(), diff --git a/studio/backend/hub/services/snapshot_progress.py b/studio/backend/hub/services/snapshot_progress.py index 1fdf05e2e5..c3db6fed7a 100644 --- a/studio/backend/hub/services/snapshot_progress.py +++ b/studio/backend/hub/services/snapshot_progress.py @@ -86,9 +86,20 @@ def _snapshot_complete_on_disk( return False if variant is None and hf_cache_scan.repo_cache_dir_has_incomplete_blobs(entry): return False - if download_manifest.has_cancel_marker(repo_type, repo_id, variant): + hub_cache = entry.parent + if download_manifest.has_cancel_marker( + repo_type, + repo_id, + variant, + hub_cache = hub_cache, + ): return False - manifest = download_manifest.read_manifest(repo_type, repo_id, variant) + manifest = download_manifest.read_manifest( + repo_type, + repo_id, + variant, + hub_cache = hub_cache, + ) if manifest is None: return False return download_manifest.verify_against_disk(manifest, snapshot_dir).ok @@ -118,6 +129,8 @@ def compute_snapshot_progress( 0, int(getattr(metadata, "completed_baseline_bytes", 0) or 0), ) + metadata_hub_cache = getattr(metadata, "hub_cache", None) + active_root = Path(metadata_hub_cache) if metadata_hub_cache else None expected_total = max(expected_bytes, 0) # Always resolve the revision's blob hashes so stale blobs from a superseded @@ -134,11 +147,17 @@ def compute_snapshot_progress( count_finalized_unscoped = variant is None readings: list[tuple[int, int, Optional[str], bool]] = [] - for entry in preferred_repo_cache_dirs( - repo_type, - repo_id, - force_active = force_active, - ): + cache_dirs = ( + preferred_repo_cache_dirs( + repo_type, + repo_id, + force_active = force_active, + active_root = active_root, + ) + if active_root is not None + else preferred_repo_cache_dirs(repo_type, repo_id, force_active = force_active) + ) + for entry in cache_dirs: completed_bytes = 0 in_progress_bytes = 0 cache_path = hf_cache_scan.resolve_hf_cache_realpath(entry) diff --git a/studio/backend/hub/tests/test_dataset_services.py b/studio/backend/hub/tests/test_dataset_services.py index 4890714cd0..6aab07cc46 100644 --- a/studio/backend/hub/tests/test_dataset_services.py +++ b/studio/backend/hub/tests/test_dataset_services.py @@ -72,57 +72,115 @@ def test_dataset_cache_scan_merges_raw_and_processed_rows(monkeypatch): assert rows[0]["partial"] is False -def test_delete_cached_dataset_attempts_all_roots_before_raising(monkeypatch): +def test_delete_cached_dataset_scopes_delete_to_selected_root(monkeypatch, tmp_path): + """A dataset present in the active cache and a previously selected cache is + deleted only from the selected root, so the other cache's copy survives.""" calls = [] - purged_state = [] + target_hub = tmp_path / "active" / "hub" + other_hub = tmp_path / "previous" / "hub" + for hub in (target_hub, other_hub): + (hub / "datasets--Org--Data").mkdir(parents = True) class _DeleteStrategy: - def __init__(self, label: str, fail: bool): + def __init__(self, label: str): self.label = label - self.fail = fail def execute(self): calls.append(self.label) - if self.fail: - raise RuntimeError(f"{self.label} failed") - class _Cache: - def __init__(self, label: str, fail: bool): - self.cache_dir = label - self.repos = [ + def _cache(label: str, hub): + return SimpleNamespace( + cache_dir = label, + repos = [ SimpleNamespace( repo_type = "dataset", repo_id = "Org/Data", + repo_path = str(hub / "datasets--Org--Data"), revisions = [SimpleNamespace(commit_hash = f"{label}-rev")], ) - ] - self.fail = fail - - def delete_revisions(self, *_revisions): - return _DeleteStrategy(self.cache_dir, self.fail) + ], + delete_revisions = lambda *_revs, _label = label: _DeleteStrategy(_label), + ) monkeypatch.setattr( cache_inventory, "_collect_hf_cache_scans", - lambda: ([_Cache("first", True), _Cache("second", False)], set()), + lambda: ([_cache("active", target_hub), _cache("previous", other_hub)], set()), ) monkeypatch.setattr( cache_inventory, "_delete_processed_dataset_cache", - lambda _repo_id: (True, []), + lambda _repo_id, **_kwargs: (False, []), ) monkeypatch.setattr( cache_inventory.download_manifest, "purge_all_state_for_repo", - lambda *_args: purged_state.append(True) or 1, + lambda *_args, **_kwargs: 0, + ) + monkeypatch.setattr( + "utils.hf_cache_settings.get_hf_cache_paths", + lambda: SimpleNamespace(hub_cache = target_hub), + ) + monkeypatch.setattr( + "hub.utils.hf_cache_state.hf_cache_roots", + lambda: [target_hub, other_hub], ) - with pytest.raises(HTTPException) as exc_info: - cache_inventory._delete_cached_dataset_blocking("Org/Data") + result = cache_inventory._delete_cached_dataset_blocking("Org/Data") - assert exc_info.value.status_code == 500 - assert calls == ["first", "second"] - assert purged_state == [] + assert result == {"status": "deleted", "repo_id": "Org/Data"} + # Only the selected (active) cache's revision is deleted; the previous + # cache's copy is never touched. + assert calls == ["active"] + assert not (target_hub / "datasets--Org--Data").exists() + assert (other_hub / "datasets--Org--Data").exists() + + +def test_delete_processed_only_dataset_accepts_processed_cache_path(monkeypatch, tmp_path): + """A processed-only dataset row sends its Arrow cache path (___ + under HF_DATASETS_CACHE), which is not a Hub datasets-- dir. The delete must + accept it and run the processed-cache delete instead of raising 400.""" + datasets_root = tmp_path / "datasets" + processed_dir = datasets_root / "Org___Data" + processed_dir.mkdir(parents = True) + + # No Hub-cache copy exists; only the processed Arrow cache holds this repo. + monkeypatch.setattr(cache_inventory, "_collect_hf_cache_scans", lambda: ([], set())) + monkeypatch.setattr(cache_inventory, "_hf_datasets_cache_roots", lambda: [datasets_root]) + processed_calls: list[str] = [] + monkeypatch.setattr( + cache_inventory, + "_delete_processed_dataset_cache", + lambda repo_id, **_kwargs: (processed_calls.append(repo_id) or True, []), + ) + + result = cache_inventory._delete_cached_dataset_blocking("Org/Data", str(processed_dir)) + + assert result == {"status": "deleted", "repo_id": "Org/Data"} + assert processed_calls == ["Org/Data"] + + +def test_delete_processed_dataset_scopes_to_selected_root(monkeypatch, tmp_path): + """A dataset processed under two HF_DATASETS_CACHE roots is deleted only from + the selected root; the copy under the other cache home survives (real delete, + not stubbed).""" + selected_root = tmp_path / "selected" / "datasets" + other_root = tmp_path / "other" / "datasets" + for root in (selected_root, other_root): + (root / "Org___Data").mkdir(parents = True) + + monkeypatch.setattr(cache_inventory, "_collect_hf_cache_scans", lambda: ([], set())) + monkeypatch.setattr( + cache_inventory, "_hf_datasets_cache_roots", lambda: [selected_root, other_root] + ) + + result = cache_inventory._delete_cached_dataset_blocking( + "Org/Data", str(selected_root / "Org___Data") + ) + + assert result == {"status": "deleted", "repo_id": "Org/Data"} + assert not (selected_root / "Org___Data").exists() # the selected copy is deleted + assert (other_root / "Org___Data").exists() # the other cache home is untouched def test_delete_cached_dataset_purges_blob_only_repo_dir(monkeypatch): @@ -139,22 +197,22 @@ def test_delete_cached_dataset_purges_blob_only_repo_dir(monkeypatch): monkeypatch.setattr( cache_inventory, "_delete_processed_dataset_cache", - lambda _repo_id: (False, []), + lambda _repo_id, **_kwargs: (False, []), ) monkeypatch.setattr( cache_inventory, "purge_repo_cache_dirs", - lambda _repo_type, repo_id: purged_dirs.append(repo_id) or True, + lambda _repo_type, repo_id, **_kwargs: purged_dirs.append(repo_id) or True, ) monkeypatch.setattr( cache_inventory, "purge_partial_repo", - lambda *_args: False, + lambda *_args, **_kwargs: False, ) monkeypatch.setattr( cache_inventory.download_manifest, "purge_all_state_for_repo", - lambda *_args: 0, + lambda *_args, **_kwargs: 0, ) result = cache_inventory._delete_cached_dataset_blocking("Org/Data") @@ -172,22 +230,22 @@ def test_delete_cached_dataset_absent_everywhere_raises_404(monkeypatch): monkeypatch.setattr( cache_inventory, "_delete_processed_dataset_cache", - lambda _repo_id: (False, []), + lambda _repo_id, **_kwargs: (False, []), ) monkeypatch.setattr( cache_inventory, "purge_repo_cache_dirs", - lambda *_args: False, + lambda *_args, **_kwargs: False, ) monkeypatch.setattr( cache_inventory, "purge_partial_repo", - lambda *_args: False, + lambda *_args, **_kwargs: False, ) monkeypatch.setattr( cache_inventory.download_manifest, "purge_all_state_for_repo", - lambda *_args: 0, + lambda *_args, **_kwargs: 0, ) with pytest.raises(HTTPException) as exc_info: diff --git a/studio/backend/hub/tests/test_download_manifest_scoping.py b/studio/backend/hub/tests/test_download_manifest_scoping.py new file mode 100644 index 0000000000..966eeaf0c2 --- /dev/null +++ b/studio/backend/hub/tests/test_download_manifest_scoping.py @@ -0,0 +1,62 @@ +# SPDX-License-Identifier: AGPL-3.0-only +# Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. See /studio/LICENSE.AGPL-3.0 + +import json +from types import SimpleNamespace + +from hub.utils import download_manifest, state_dir + + +def _write_manifest(path, payload): + path.parent.mkdir(parents = True, exist_ok = True) + path.write_text(json.dumps(payload), encoding = "utf-8") + + +def test_purge_state_preserves_active_legacy_when_deleting_inactive_cache(monkeypatch, tmp_path): + """A scoped delete of an inactive cache must not erase the unscoped legacy + state, which _legacy_state_applies attributes to the active cache.""" + active = tmp_path / "active" / "hub" + previous = tmp_path / "previous" / "hub" + for path in (active, previous): + path.mkdir(parents = True) + + monkeypatch.setattr(state_dir, "cache_root", lambda: tmp_path / "state") + monkeypatch.setattr( + "utils.hf_cache_settings.get_hf_cache_paths", + lambda: SimpleNamespace(hub_cache = str(active)), + ) + + # Unowned legacy manifest -> belongs to the active cache. + legacy = state_dir.manifest_path("model", "Org/Model") + _write_manifest(legacy, {"version": 1}) + # The inactive cache's own scoped copy is the one being deleted. + scoped = state_dir.manifest_path("model", "Org/Model", hub_cache = str(previous)) + _write_manifest(scoped, {"version": 1, "hub_cache": str(previous)}) + + removed = download_manifest.purge_state("model", "Org/Model", hub_cache = str(previous)) + + assert removed is True + assert not scoped.is_file() # the inactive cache's copy is gone + assert legacy.is_file() # the active cache's legacy state survives + + +def test_purge_state_removes_legacy_owned_by_the_deleted_cache(monkeypatch, tmp_path): + """A legacy file that recorded the deleted cache as its owner is purged.""" + active = tmp_path / "active" / "hub" + previous = tmp_path / "previous" / "hub" + for path in (active, previous): + path.mkdir(parents = True) + + monkeypatch.setattr(state_dir, "cache_root", lambda: tmp_path / "state") + monkeypatch.setattr( + "utils.hf_cache_settings.get_hf_cache_paths", + lambda: SimpleNamespace(hub_cache = str(active)), + ) + + legacy = state_dir.manifest_path("model", "Org/Model") + _write_manifest(legacy, {"version": 1, "hub_cache": str(previous)}) + + removed = download_manifest.purge_state("model", "Org/Model", hub_cache = str(previous)) + + assert removed is True + assert not legacy.is_file() # owned by the deleted cache -> purged diff --git a/studio/backend/hub/tests/test_empty_variant_folder.py b/studio/backend/hub/tests/test_empty_variant_folder.py index 33bf6c6819..3ed8e69e0d 100644 --- a/studio/backend/hub/tests/test_empty_variant_folder.py +++ b/studio/backend/hub/tests/test_empty_variant_folder.py @@ -120,10 +120,16 @@ def _force_compute_to_raise(monkeypatch): monkeypatch.setattr(gguf_variants, "list_gguf_variants", _boom, raising = False) monkeypatch.setattr( - gguf_variants, "list_gguf_variants_from_hf_cache", lambda repo_id: None, raising = False + gguf_variants, + "list_gguf_variants_from_hf_cache", + lambda repo_id, root = None: None, + raising = False, ) monkeypatch.setattr( - gguf_variants, "list_partial_gguf_variants_from_state", lambda repo_id: None, raising = False + gguf_variants, + "list_partial_gguf_variants_from_state", + lambda repo_id, hub_cache = None: None, + raising = False, ) @@ -133,7 +139,11 @@ def test_get_variants_surfaces_cleanable_when_metadata_fails(monkeypatch): import asyncio _force_compute_to_raise(monkeypatch) - monkeypatch.setattr(gguf_variants, "list_empty_gguf_variant_dirs", lambda repo_id: {"UD-IQ1_S"}) + monkeypatch.setattr( + gguf_variants, + "list_empty_gguf_variant_dirs", + lambda repo_id, root = None: {"UD-IQ1_S"}, + ) resp = asyncio.run( gguf_variants.get_gguf_variants_response( @@ -152,7 +162,11 @@ def test_get_variants_reraises_when_no_cleanable(monkeypatch): from fastapi import HTTPException _force_compute_to_raise(monkeypatch) - monkeypatch.setattr(gguf_variants, "list_empty_gguf_variant_dirs", lambda repo_id: set()) + monkeypatch.setattr( + gguf_variants, + "list_empty_gguf_variant_dirs", + lambda repo_id, root = None: set(), + ) try: asyncio.run( diff --git a/studio/backend/hub/tests/test_model_services.py b/studio/backend/hub/tests/test_model_services.py index 693d945ee1..fa5862a13c 100644 --- a/studio/backend/hub/tests/test_model_services.py +++ b/studio/backend/hub/tests/test_model_services.py @@ -2,6 +2,7 @@ # Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. See /studio/LICENSE.AGPL-3.0 import asyncio +import json import sys from pathlib import Path from types import SimpleNamespace @@ -102,6 +103,236 @@ def test_big_endian_detection_ignores_model_name_be_token(): ) +def _cached_model_row(tmp_path: Path, *, partial: bool, active_cache: bool | None, size_bytes: int): + path = tmp_path / f"cache-{active_cache}-{partial}-{size_bytes}" + return model_common._local_model_info( + scan_path = path, + load_path = path, + source = "hf_cache", + model_format = "safetensors", + model_id = "Org/Model", + partial = partial, + active_cache = active_cache, + size_bytes = size_bytes, + ) + + +@pytest.mark.parametrize("reverse", [False, True]) +def test_local_inventory_prefers_complete_previous_cache_copy(tmp_path, reverse): + active_partial = _cached_model_row( + tmp_path, + partial = True, + active_cache = True, + size_bytes = 20, + ) + previous_complete = _cached_model_row( + tmp_path, + partial = False, + active_cache = False, + size_bytes = 10, + ) + rows = [active_partial, previous_complete] + if reverse: + rows.reverse() + + result = local_inventory._dedupe_local_models(rows) + + assert result == [previous_complete] + + +def test_local_inventory_compares_all_non_active_cache_copies(tmp_path): + inactive_partial = _cached_model_row( + tmp_path, + partial = True, + active_cache = False, + size_bytes = 20, + ) + custom_complete = _cached_model_row( + tmp_path, + partial = False, + active_cache = None, + size_bytes = 10, + ) + + assert local_inventory._dedupe_local_models([inactive_partial, custom_complete]) == [ + custom_complete + ] + + +def test_local_inventory_prefers_active_cache_when_copies_are_equally_complete(tmp_path): + previous = _cached_model_row( + tmp_path, + partial = False, + active_cache = False, + size_bytes = 20, + ) + active = _cached_model_row( + tmp_path, + partial = False, + active_cache = True, + size_bytes = 10, + ) + + assert local_inventory._dedupe_local_models([previous, active]) == [active] + + +def test_loaded_repo_match_accepts_previous_cache_snapshot_path(monkeypatch, tmp_path): + repo_dir = tmp_path / "old-hub" / "models--Org--Model" + snapshot = repo_dir / "snapshots" / "revision" + snapshot.mkdir(parents = True) + monkeypatch.setattr(deletion, "iter_repo_cache_dirs", lambda *_args: iter([repo_dir])) + + assert deletion._loaded_id_matches_repo(str(snapshot), "Org/Model") is True + assert deletion._loaded_id_matches_repo(str(snapshot / "model.gguf"), "Org/Model") is True + assert deletion._loaded_id_matches_repo(str(tmp_path / "other"), "Org/Model") is False + + +def test_cached_inventory_loads_previous_cache_copy_by_snapshot(monkeypatch, tmp_path): + active_hub = tmp_path / "active-hub" + previous_repo = tmp_path / "previous-hub" / "models--Org--Model" + snapshot = previous_repo / "snapshots" / "revision" + snapshot.mkdir(parents = True) + monkeypatch.setattr( + "utils.hf_cache_settings.get_hf_cache_paths", + lambda: SimpleNamespace(hub_cache = active_hub), + ) + + fields = cache_inventory._cache_inventory_fields( + "Org/Model", + "safetensors", + repo_path = previous_repo, + snapshot_path = snapshot, + ) + + assert fields["load_id"] == str(snapshot) + + +def test_cached_inventory_keeps_repo_id_for_active_cache(monkeypatch, tmp_path): + active_hub = tmp_path / "active-hub" + active_repo = active_hub / "models--Org--Model" + monkeypatch.setattr( + "utils.hf_cache_settings.get_hf_cache_paths", + lambda: SimpleNamespace(hub_cache = active_hub), + ) + + fields = cache_inventory._cache_inventory_fields( + "Org/Model", + "safetensors", + repo_path = active_repo, + ) + + assert fields["load_id"] == "Org/Model" + + +def test_cached_inventory_prefers_active_copy_when_completeness_matches(): + previous = {"partial": False, "active_cache": False, "size_bytes": 200} + active = {"partial": False, "active_cache": True, "size_bytes": 100} + + assert cache_inventory._prefer_cache_row(active, previous) is True + assert cache_inventory._prefer_cache_row(previous, active) is False + + +def test_cached_inventory_prefers_complete_copy_before_active_cache(): + previous = {"partial": False, "active_cache": False, "size_bytes": 100} + active_partial = {"partial": True, "active_cache": True, "size_bytes": 200} + + assert cache_inventory._prefer_cache_row(previous, active_partial) is True + assert cache_inventory._prefer_cache_row(active_partial, previous) is False + + +def test_inventory_scans_every_dynamic_cache_root(monkeypatch, tmp_path): + first = tmp_path / "first-hub" + second = tmp_path / "second-hub" + unreadable = tmp_path / "unreadable-hub" + first.mkdir() + second.mkdir() + unreadable.mkdir() + scanned = [] + + monkeypatch.setattr( + inventory_scan, + "hf_cache_roots", + lambda: [first, unreadable, second], + ) + + def scan_cache(cache_dir): + path = Path(cache_dir) + scanned.append(path) + if path == unreadable: + raise PermissionError("unreadable") + return SimpleNamespace(cache_dir = cache_dir) + + monkeypatch.setattr("huggingface_hub.scan_cache_dir", scan_cache) + + result = inventory_scan._compute_all_hf_cache_scans() + + assert scanned == [first, unreadable, second] + assert [Path(scan.cache_dir) for scan in result] == [first, second] + + +def test_inventory_applies_download_state_to_its_owning_cache(monkeypatch, tmp_path): + state_root = tmp_path / "state" + cache_a = tmp_path / "cache-a" + cache_b = tmp_path / "cache-b" + repo_id = "Org/Model" + repo_name = "models--Org--Model" + repo_a = cache_a / repo_name + repo_b = cache_b / repo_name + snapshot_a = repo_a / "snapshots" / "revision" + snapshot_b = repo_b / "snapshots" / "revision" + snapshot_a.mkdir(parents = True) + snapshot_b.mkdir(parents = True) + (snapshot_a / "config.json").write_bytes(b"x") + (snapshot_b / "config.json").write_bytes(b"xx") + + monkeypatch.setattr(state_dir, "cache_root", lambda: state_root) + monkeypatch.setattr( + "utils.hf_cache_settings.get_hf_cache_paths", + lambda: SimpleNamespace(hub_cache = cache_b), + ) + assert download_manifest.write_manifest( + "model", + repo_id, + None, + [download_manifest.ExpectedFile(path = "config.json", size = 2)], + "http", + hub_cache = cache_a, + ) + + assert inventory_scan.is_snapshot_partial("model", repo_id, repo_a) is True + assert inventory_scan.is_snapshot_partial("model", repo_id, repo_b) is False + assert inventory_scan.partial_transport_for("model", repo_id, None, repo_a) == "http" + assert inventory_scan.partial_transport_for("model", repo_id, None, repo_b) is None + + +def test_inventory_scopes_cancel_markers_to_their_owning_cache(monkeypatch, tmp_path): + state_root = tmp_path / "state" + cache_a = tmp_path / "cache-a" + cache_b = tmp_path / "cache-b" + repo_id = "Org/Model" + repo_name = "models--Org--Model" + repo_a = cache_a / repo_name + repo_b = cache_b / repo_name + repo_a.mkdir(parents = True) + repo_b.mkdir(parents = True) + + monkeypatch.setattr(state_dir, "cache_root", lambda: state_root) + monkeypatch.setattr( + "utils.hf_cache_settings.get_hf_cache_paths", + lambda: SimpleNamespace(hub_cache = cache_b), + ) + assert download_manifest.write_cancel_marker( + "model", + repo_id, + "Q4_K_M", + "xet", + hub_cache = cache_a, + ) + + assert inventory_scan.is_variant_partial(repo_id, "Q4_K_M", repo_cache_dir = repo_a) is True + assert inventory_scan.is_variant_partial(repo_id, "Q4_K_M", repo_cache_dir = repo_b) is False + + def test_list_local_gguf_variants_skips_big_endian_sibling(tmp_path): (tmp_path / "model-Q4_K_M-be.gguf").write_bytes(b"x" * 100) (tmp_path / "model-Q4_K_M.gguf").write_bytes(b"y" * 10) @@ -163,8 +394,19 @@ def test_download_state_bounds_long_repo_variant_filenames(monkeypatch, tmp_path "http", ) - marker_path = state_dir.marker_path("model", repo_id, variant) - manifest_path = state_dir.manifest_path("model", repo_id, variant) + hub_cache = download_manifest._canonical_hub_cache() + marker_path = state_dir.marker_path( + "model", + repo_id, + variant, + hub_cache = hub_cache, + ) + manifest_path = state_dir.manifest_path( + "model", + repo_id, + variant, + hub_cache = hub_cache, + ) assert marker_path is not None assert manifest_path is not None @@ -181,6 +423,97 @@ def test_download_state_bounds_long_repo_variant_filenames(monkeypatch, tmp_path ] +def test_download_state_isolated_across_hub_cache_switches(monkeypatch, tmp_path): + monkeypatch.setattr(state_dir, "cache_root", lambda: tmp_path) + cache_a = tmp_path / "cache-a" + cache_b = tmp_path / "cache-b" + selected = SimpleNamespace(hub_cache = cache_a) + + from utils import hf_cache_settings + + monkeypatch.setattr(hf_cache_settings, "get_hf_cache_paths", lambda: selected) + expected_a = [download_manifest.ExpectedFile(path = "a.gguf", size = 1)] + expected_b = [download_manifest.ExpectedFile(path = "b.gguf", size = 2)] + + assert download_manifest.write_manifest("model", "Owner/Repo", "Q4_K_M", expected_a) + assert download_manifest.write_cancel_marker("model", "Owner/Repo", "Q4_K_M", "http") + + selected.hub_cache = cache_b + assert download_manifest.write_manifest("model", "Owner/Repo", "Q4_K_M", expected_b) + + manifest_b = download_manifest.read_manifest("model", "Owner/Repo", "Q4_K_M") + manifest_a = download_manifest.read_manifest( + "model", + "Owner/Repo", + "Q4_K_M", + hub_cache = cache_a, + ) + + assert manifest_b is not None and manifest_b.expected_files == tuple(expected_b) + assert manifest_a is not None and manifest_a.expected_files == tuple(expected_a) + assert not download_manifest.has_cancel_marker("model", "Owner/Repo", "Q4_K_M") + assert download_manifest.has_cancel_marker( + "model", + "Owner/Repo", + "Q4_K_M", + hub_cache = cache_a, + ) + assert len(list((tmp_path / "hub-state" / "manifests").rglob("*.json"))) == 2 + + +def test_legacy_unscoped_download_state_falls_back_only_for_selected_cache(monkeypatch, tmp_path): + monkeypatch.setattr(state_dir, "cache_root", lambda: tmp_path) + cache_a = tmp_path / "cache-a" + cache_b = tmp_path / "cache-b" + monkeypatch.setattr( + "utils.hf_cache_settings.get_hf_cache_paths", + lambda: SimpleNamespace(hub_cache = cache_a), + ) + manifest = state_dir.manifest_path("model", "Owner/Repo", "Q4_K_M") + marker = state_dir.marker_path("model", "Owner/Repo", "Q4_K_M") + assert manifest is not None and marker is not None + manifest.write_text( + json.dumps( + { + "version": 1, + "repo_id": "Owner/Repo", + "variant": "Q4_K_M", + "expected_files": [{"path": "model.gguf", "size": 10}], + "transport": "http", + } + ), + encoding = "utf-8", + ) + marker.write_text( + json.dumps({"version": 1, "repo_id": "Owner/Repo", "variant": "Q4_K_M"}), + encoding = "utf-8", + ) + + assert download_manifest.read_manifest("model", "Owner/Repo", "Q4_K_M") is not None + assert download_manifest.has_cancel_marker("model", "Owner/Repo", "Q4_K_M") + assert list(download_manifest.iter_variant_manifests("model", "Owner/Repo")) == [ + ("Q4_K_M", manifest) + ] + assert list(download_manifest.iter_variant_markers("model", "Owner/Repo")) == [ + ("Q4_K_M", marker) + ] + assert ( + download_manifest.read_manifest( + "model", + "Owner/Repo", + "Q4_K_M", + hub_cache = cache_b, + ) + is None + ) + assert not download_manifest.has_cancel_marker( + "model", + "Owner/Repo", + "Q4_K_M", + hub_cache = cache_b, + ) + + class _RecordingLogger: def __init__(self): self.warnings = [] @@ -416,8 +749,15 @@ def test_cached_gguf_scan_includes_variant_state_without_completed_gguf(monkeypa "Q4_K_M", [download_manifest.ExpectedFile(path = "model-Q4_K_M.gguf", size = 4096)], "http", + hub_cache = repo_path.parent, + ) + assert download_manifest.write_cancel_marker( + "model", + "Org/PartialGguf", + "Q4_K_M", + "http", + hub_cache = repo_path.parent, ) - assert download_manifest.write_cancel_marker("model", "Org/PartialGguf", "Q4_K_M", "http") monkeypatch.setattr( cache_inventory, "all_hf_cache_scans", @@ -484,6 +824,7 @@ def test_cached_gguf_scan_keeps_infra_repo_with_user_downloaded_variant(monkeypa "Q8_0", [download_manifest.ExpectedFile(path = "bge-small-en-v1.5-Q8_0.gguf", size = 35_000_000)], "http", + hub_cache = Path(embedder.repo_path).parent, ) monkeypatch.setattr( cache_inventory, @@ -1206,6 +1547,7 @@ def test_gguf_progress_counts_completed_mmproj_with_expected_bytes(monkeypatch, ), ], "http", + hub_cache = entry.parent, ) requirement = gguf_variants._GgufVariantRequirement( @@ -1292,6 +1634,7 @@ def test_gguf_progress_subtracts_new_job_completed_baseline(monkeypatch, tmp_pat ), ], "http", + hub_cache = entry.parent, ) requirement = gguf_variants._GgufVariantRequirement( @@ -1462,6 +1805,7 @@ def test_gguf_progress_complete_on_disk_ignores_full_baseline(monkeypatch, tmp_p ), ], "http", + hub_cache = entry.parent, ) requirement = gguf_variants._GgufVariantRequirement( @@ -1861,8 +2205,15 @@ def test_hf_cache_scan_uses_gguf_partial_row_for_variant_state(monkeypatch, tmp_ "Q4_K_M", [download_manifest.ExpectedFile(path = "model-Q4_K_M.gguf", size = 8192)], "http", + hub_cache = cache_dir, + ) + assert download_manifest.write_cancel_marker( + "model", + "Org/PartialGguf", + "Q4_K_M", + "http", + hub_cache = cache_dir, ) - assert download_manifest.write_cancel_marker("model", "Org/PartialGguf", "Q4_K_M", "http") monkeypatch.setattr(local_inventory, "_classify_local_path", lambda *_args, **_kwargs: []) monkeypatch.setattr( local_inventory.hf_cache_scan, @@ -2117,7 +2468,7 @@ def test_gguf_variants_partial_marker_overrides_size_only_downloaded(monkeypatch monkeypatch.setattr( gguf_variants, "iter_hf_cache_snapshots", - lambda _repo_id: [snapshot], + lambda _repo_id, root = None: [snapshot], ) monkeypatch.setattr( gguf_variants, @@ -2136,6 +2487,70 @@ def test_gguf_variants_partial_marker_overrides_size_only_downloaded(monkeypatch assert result.variants[0].partial is True +def test_gguf_variants_scopes_partial_state_to_requested_cache(monkeypatch, tmp_path): + async def _run_inline(fn, *args, **kwargs): + return fn(*args, **kwargs) + + repo_id = "Org/SharedRepo" + repo_name = "models--Org--SharedRepo" + cache_a = tmp_path / "cache-a" + cache_b = tmp_path / "cache-b" + repo_a = cache_a / repo_name + snapshot_a = repo_a / "snapshots" / "revision" + snapshot_a.mkdir(parents = True) + (snapshot_a / "model-Q8_0.gguf").write_bytes(b"complete") + blobs_b = cache_b / repo_name / "blobs" + blobs_b.mkdir(parents = True) + (blobs_b / "q8-hash.incomplete").write_bytes(b"partial") + + monkeypatch.setattr(state_dir, "cache_root", lambda: tmp_path / "state") + monkeypatch.setattr(gguf_variants.asyncio, "to_thread", _run_inline) + monkeypatch.setattr( + "utils.hf_cache_settings.get_hf_cache_paths", + lambda: SimpleNamespace(hub_cache = cache_b), + ) + assert download_manifest.write_cancel_marker( + "model", + repo_id, + "Q8_0", + "http", + hub_cache = cache_b, + ) + monkeypatch.setattr( + gguf_variants, + "list_gguf_variants", + lambda *_args, **_kwargs: ( + [ + SimpleNamespace( + filename = "model-Q8_0.gguf", + quant = "Q8_0", + display_label = None, + size_bytes = 8, + ) + ], + False, + [ + SimpleNamespace( + rfilename = "model-Q8_0.gguf", + size = 8, + lfs = SimpleNamespace(sha256 = "q8-hash"), + ) + ], + ), + ) + monkeypatch.setattr(cache_inventory, "all_hf_cache_scans", lambda: []) + + result = asyncio.run( + gguf_variants.get_gguf_variants_response( + repo_id, + local_path = str(repo_a), + ) + ) + + assert result.variants[0].downloaded is True + assert result.variants[0].partial is False + + def test_download_registry_repo_keys_are_case_insensitive(): registry = download_registry.DownloadRegistry() @@ -2444,6 +2859,34 @@ def test_prepare_cache_for_transport_purges_only_requested_hashes(monkeypatch, t assert (blobs / "shared-mmproj.incomplete").exists() +def test_prepare_cache_for_transport_uses_captured_root(monkeypatch, tmp_path): + cache_a = tmp_path / "cache-a" + cache_b = tmp_path / "cache-b" + repo_name = "models--Org--Repo" + partial_a = cache_a / repo_name / "blobs" / "blob.incomplete" + partial_b = cache_b / repo_name / "blobs" / "blob.incomplete" + partial_a.parent.mkdir(parents = True) + partial_b.parent.mkdir(parents = True) + partial_a.write_bytes(b"a") + partial_b.write_bytes(b"b") + monkeypatch.setattr( + download_registry, + "hf_cache_root", + lambda create = False, root = None: root or cache_b, + ) + + purged = download_registry.prepare_cache_for_transport( + "model", + "Org/Repo", + download_registry.TRANSPORT_HTTP, + root = cache_a, + ) + + assert purged == 1 + assert not partial_a.exists() + assert partial_b.exists() + + def _vision_cache_root(monkeypatch, tmp_path): root = tmp_path / "hub" blobs = root / "models--Org--Vision" / "blobs" @@ -2802,6 +3245,47 @@ def test_shutdown_skips_marker_for_worker_that_exits_cleanly(monkeypatch): assert markers == ["Org/Cut"] +def test_orphan_reaper_uses_worker_cache_root_after_setting_changes(monkeypatch, tmp_path): + workers = tmp_path / "workers" + workers.mkdir() + cache_a = tmp_path / "cache-a" / "hub" + cache_b = tmp_path / "cache-b" / "hub" + partial = cache_a / "models--Org--Model" / "blobs" / "abc.incomplete" + partial.parent.mkdir(parents = True) + partial.write_bytes(b"partial") + cache_b.mkdir(parents = True) + monkeypatch.setattr(state_dir, "workers_dir", lambda: workers) + monkeypatch.setattr(download_registry, "_process_alive", lambda _pid: False) + monkeypatch.setattr( + "utils.hf_cache_settings.get_hf_cache_paths", + lambda: SimpleNamespace(hub_cache = cache_b), + ) + markers = [] + monkeypatch.setattr( + download_registry, + "persist_cancel_marker", + lambda *args, **kwargs: markers.append(args), + ) + metadata = download_registry.DownloadMetadata( + repo_type = "model", + repo_id = "Org/Model", + variant = None, + transport = download_registry.TRANSPORT_HTTP, + hub_cache = str(cache_a), + xet_cache = str(tmp_path / "cache-a" / "xet"), + ) + download_registry.write_worker_breadcrumb("org/model", 1234, metadata) + [breadcrumb] = list(workers.iterdir()) + payload = json.loads(breadcrumb.read_text(encoding = "utf-8")) + assert payload["hub_cache"] == str(cache_a) + assert payload["xet_cache"] == str(tmp_path / "cache-a" / "xet") + + download_registry.reap_orphan_workers() + + assert markers == [("model", "Org/Model", None, "http")] + assert list(workers.iterdir()) == [] + + def test_model_claim_register_cancel_uses_registry_marker_owner(monkeypatch): killed = [] @@ -3125,12 +3609,19 @@ def _build_variant_cache_repo(repo_dir, blob_specs, snapshot_links): return repo -def _patch_variant_delete_side_effects(monkeypatch): +def _patch_variant_delete_side_effects(monkeypatch, hub_cache = None): monkeypatch.setattr( deletion.download_manifest, "purge_state", lambda *_args, **_kwargs: False, ) + # The repo under test lives in this cache; make it the active one so the + # delete scopes to it (default target root is the active hub cache). + if hub_cache is not None: + monkeypatch.setattr( + "utils.hf_cache_settings.get_hf_cache_paths", + lambda: SimpleNamespace(hub_cache = hub_cache), + ) def test_snapshot_progress_filters_stale_blobs(monkeypatch, tmp_path): @@ -3308,7 +3799,7 @@ def test_delete_variant_keeps_blob_shared_with_other_snapshot(monkeypatch, tmp_p "all_hf_cache_scans", lambda: [SimpleNamespace(repos = [repo])], ) - _patch_variant_delete_side_effects(monkeypatch) + _patch_variant_delete_side_effects(monkeypatch, tmp_path) result = deletion._delete_cached_model_blocking("Org/Repo-GGUF", "Q4_K_M", None) @@ -3335,7 +3826,7 @@ def test_delete_variant_unlinks_unshared_blob(monkeypatch, tmp_path): "all_hf_cache_scans", lambda: [SimpleNamespace(repos = [repo])], ) - _patch_variant_delete_side_effects(monkeypatch) + _patch_variant_delete_side_effects(monkeypatch, tmp_path) result = deletion._delete_cached_model_blocking("Org/Repo-GGUF", "Q4_K_M", None) @@ -3361,7 +3852,7 @@ def test_delete_variant_surfaces_locked_file_as_conflict(monkeypatch, tmp_path): "all_hf_cache_scans", lambda: [SimpleNamespace(repos = [repo])], ) - _patch_variant_delete_side_effects(monkeypatch) + _patch_variant_delete_side_effects(monkeypatch, tmp_path) real_unlink = Path.unlink diff --git a/studio/backend/hub/utils/download_manifest.py b/studio/backend/hub/utils/download_manifest.py index 5366689296..ac0ccb5490 100644 --- a/studio/backend/hub/utils/download_manifest.py +++ b/studio/backend/hub/utils/download_manifest.py @@ -77,6 +77,7 @@ class Manifest: started_at: str expected_files: tuple[ExpectedFile, ...] transport: Optional[str] = None + hub_cache: Optional[str] = None @dataclass(frozen = True) @@ -86,6 +87,78 @@ class VerifyResult: size_mismatched: tuple[str, ...] +def _canonical_hub_cache(hub_cache: Optional[str | Path] = None) -> Optional[str]: + if hub_cache is None: + try: + from utils.hf_cache_settings import get_hf_cache_paths + hub_cache = get_hf_cache_paths().hub_cache + except Exception: + return None + try: + return str(Path(hub_cache).expanduser().resolve(strict = False)) + except (OSError, RuntimeError, ValueError): + return str(hub_cache) + + +def _read_state_payload(path: Path) -> Optional[dict]: + try: + data = json.loads(path.read_text(encoding = "utf-8")) + except (OSError, ValueError) as exc: + logger.debug("Could not read Hub state %s: %s", path, exc) + return None + return data if isinstance(data, dict) else None + + +def _legacy_state_applies( + path: Path, + requested_hub_cache: Optional[str], + *, + fail_closed: bool = False, +) -> bool: + """Whether an old unscoped state file belongs to the requested cache. + + Transitional files that recorded their cache keep that ownership. Older + files with no ownership can only be attributed to the currently selected + cache, which matches the single-cache behavior under which they were + written without leaking them into remembered inactive caches. + """ + data = _read_state_payload(path) + if data is not None: + recorded = data.get("hub_cache") + if isinstance(recorded, str) and recorded: + return _canonical_hub_cache(recorded) == requested_hub_cache + elif not fail_closed: + return False + return requested_hub_cache == _canonical_hub_cache() + + +def _state_read_path( + path_factory, + repo_type: RepoType, + repo_id: str, + variant: Optional[str], + hub_cache: Optional[str | Path], + *, + fail_closed: bool = False, +) -> Optional[Path]: + requested = _canonical_hub_cache(hub_cache) + scoped = path_factory(repo_type, repo_id, variant, hub_cache = requested) + try: + if scoped is not None and scoped.is_file(): + return scoped + except OSError: + pass + legacy = path_factory(repo_type, repo_id, variant) + if legacy is None or legacy == scoped: + return None + try: + if not legacy.is_file(): + return None + except OSError: + return None + return legacy if _legacy_state_applies(legacy, requested, fail_closed = fail_closed) else None + + def _atomic_write_json(path: Path, payload: dict) -> bool: # Per-write uuid suffix so a concurrent caller or a stale tmp from a # previous crash cannot collide with the in-flight write. @@ -124,6 +197,8 @@ def write_manifest( variant: Optional[str], expected_files: Sequence[ExpectedFile], transport: Optional[str] = None, + *, + hub_cache: Optional[str | Path] = None, ) -> bool: """Write/overwrite the manifest for this triple. Best-effort. @@ -131,7 +206,13 @@ def write_manifest( worst-case fallback is the pre-fix scanner behavior (one missed partial detection), which is no regression. """ - path = manifest_path(repo_type, repo_id, variant) + recorded_hub_cache = _canonical_hub_cache(hub_cache) + path = manifest_path( + repo_type, + repo_id, + variant, + hub_cache = recorded_hub_cache, + ) if path is None: return False payload = { @@ -149,6 +230,7 @@ def write_manifest( for f in expected_files ], "transport": transport, + "hub_cache": recorded_hub_cache, } return _atomic_write_json(path, payload) @@ -157,6 +239,8 @@ def read_manifest( repo_type: RepoType, repo_id: str, variant: Optional[str] = None, + *, + hub_cache: Optional[str | Path] = None, ) -> Optional[Manifest]: """Return the manifest if present and parseable; ``None`` otherwise. @@ -171,15 +255,17 @@ def read_manifest( ``_MANIFEST_VERSION`` and widen this check) or live under a different filename, so an incompatible payload can never mis-classify rows. """ - path = manifest_path(repo_type, repo_id, variant) + path = _state_read_path( + manifest_path, + repo_type, + repo_id, + variant, + hub_cache, + ) if path is None or not path.is_file(): return None - try: - data = json.loads(path.read_text(encoding = "utf-8")) - except (OSError, ValueError) as exc: - logger.debug("Could not read manifest %s: %s", path, exc) - return None - if not isinstance(data, dict): + data = _read_state_payload(path) + if data is None: return None if data.get("version") != _MANIFEST_VERSION: logger.debug( @@ -216,6 +302,7 @@ def read_manifest( started_at = str(data.get("started_at", "")), expected_files = tuple(expected), transport = transport if transport in ("http", "xet") else None, + hub_cache = data.get("hub_cache") if isinstance(data.get("hub_cache"), str) else None, ) @@ -289,6 +376,8 @@ def write_cancel_marker( repo_id: str, variant: Optional[str] = None, transport: Optional[str] = None, + *, + hub_cache: Optional[str | Path] = None, ) -> bool: """Record that this triple was cancelled. Idempotent across repeated cancels. @@ -296,7 +385,13 @@ def write_cancel_marker( inventory rows so the UI labels HTTP retries as continuable and XET retries as full redownloads. None is accepted for forward-compat. """ - path = marker_path(repo_type, repo_id, variant) + recorded_hub_cache = _canonical_hub_cache(hub_cache) + path = marker_path( + repo_type, + repo_id, + variant, + hub_cache = recorded_hub_cache, + ) if path is None: return False payload = { @@ -306,6 +401,7 @@ def write_cancel_marker( "variant": variant, "transport": transport, "cancelled_at": datetime.now(timezone.utc).isoformat(), + "hub_cache": recorded_hub_cache, } return _atomic_write_json(path, payload) @@ -314,6 +410,8 @@ def read_cancel_marker_transport( repo_type: RepoType, repo_id: str, variant: Optional[str] = None, + *, + hub_cache: Optional[str | Path] = None, ) -> Optional[str]: """Return the transport recorded in the cancel marker, or ``None`` if no marker exists or it is unreadable. @@ -330,15 +428,17 @@ def read_cancel_marker_transport( ``None`` keeps the neutral "Retry" label. * Unknown future versions โ†’ ``None`` (unknown layout, unknown transport). """ - path = marker_path(repo_type, repo_id, variant) + path = _state_read_path( + marker_path, + repo_type, + repo_id, + variant, + hub_cache, + ) if path is None or not path.is_file(): return None - try: - data = json.loads(path.read_text(encoding = "utf-8")) - except (OSError, ValueError) as exc: - logger.debug("Could not read cancel marker %s: %s", path, exc) - return None - if not isinstance(data, dict): + data = _read_state_payload(path) + if data is None: return None version = data.get("version") if version == _LEGACY_MARKER_VERSION: @@ -351,10 +451,30 @@ def read_cancel_marker_transport( return None +def _all_matching_state_paths( + parent: Optional[Path], repo_type: RepoType, repo_id: str, variant: Optional[str] +) -> tuple[Path, ...]: + if parent is None: + return () + legacy_path = ( + manifest_path(repo_type, repo_id, variant) + if parent.name == "manifests" + else marker_path(repo_type, repo_id, variant) + ) + if legacy_path is None: + return () + try: + return tuple(path for path in parent.rglob(legacy_path.name) if path.is_file()) + except OSError: + return () + + def clear_cancel_marker( repo_type: RepoType, repo_id: str, variant: Optional[str] = None, + *, + hub_cache: Optional[str | Path] = None, ) -> None: """Remove the cancel marker for this triple if present. @@ -362,31 +482,48 @@ def clear_cancel_marker( download-start (a fresh attempt supersedes prior cancel state) and again at successful completion (cleans up if the start clear failed). """ - path = marker_path(repo_type, repo_id, variant) - if path is None: - return - try: - path.unlink(missing_ok = True) - except OSError as exc: - logger.debug("Could not clear cancel marker %s: %s", path, exc) + requested = _canonical_hub_cache(hub_cache) + path = marker_path( + repo_type, + repo_id, + variant, + hub_cache = requested, + ) + legacy = marker_path(repo_type, repo_id, variant) + paths = [path] + if ( + legacy is not None + and legacy != path + and _legacy_state_applies(legacy, requested, fail_closed = True) + ): + paths.append(legacy) + for target in paths: + if target is None: + continue + try: + target.unlink(missing_ok = True) + except OSError as exc: + logger.debug("Could not clear cancel marker %s: %s", target, exc) def has_cancel_marker( repo_type: RepoType, repo_id: str, variant: Optional[str] = None, + *, + hub_cache: Optional[str | Path] = None, ) -> bool: - """File-existence check only. Body is never read. - - Fail-closed: a corrupt marker still returns ``True`` because the - file's existence is the signal (the user once cancelled this - triple, even if the body is unreadable). - """ - path = marker_path(repo_type, repo_id, variant) - if path is None: - return False + """Return whether a cancel marker applies to the selected cache.""" + path = _state_read_path( + marker_path, + repo_type, + repo_id, + variant, + hub_cache, + fail_closed = True, + ) try: - return path.is_file() + return path is not None and path.is_file() except OSError: return False @@ -395,48 +532,124 @@ def delete_manifest( repo_type: RepoType, repo_id: str, variant: Optional[str] = None, + *, + hub_cache: Optional[str | Path] = None, ) -> bool: - path = manifest_path(repo_type, repo_id, variant) - if path is None: - return False - try: - if not path.is_file(): - return False - path.unlink() - return True - except OSError as exc: - logger.debug("Could not delete manifest %s: %s", path, exc) - return False + requested = _canonical_hub_cache(hub_cache) + path = manifest_path( + repo_type, + repo_id, + variant, + hub_cache = requested, + ) + legacy = manifest_path(repo_type, repo_id, variant) + paths = [path] + if legacy is not None and legacy != path and _legacy_state_applies(legacy, requested): + paths.append(legacy) + removed = False + for target in paths: + if target is None: + continue + try: + if target.is_file(): + target.unlink() + removed = True + except OSError as exc: + logger.debug("Could not delete manifest %s: %s", target, exc) + return removed def purge_state( repo_type: RepoType, repo_id: str, variant: Optional[str] = None, + *, + hub_cache: Optional[str | Path] = None, ) -> bool: """Remove manifest + cancel marker for this triple. Returns ``True`` - when anything was present on disk before the call. Idempotent.""" - marker_existed = has_cancel_marker(repo_type, repo_id, variant) - manifest_removed = delete_manifest(repo_type, repo_id, variant) - clear_cancel_marker(repo_type, repo_id, variant) - return marker_existed or manifest_removed + when anything was present on disk before the call. Idempotent. + + With ``hub_cache`` set, only that cache's scoped state (plus any legacy + unscoped file that belongs to it) is removed, so purging one cache's copy + never clears another cache's resumable/cancel state.""" + if hub_cache is None: + paths = ( + *_all_matching_state_paths(manifests_dir(), repo_type, repo_id, variant), + *_all_matching_state_paths(cancelled_dir(), repo_type, repo_id, variant), + ) + else: + requested = _canonical_hub_cache(hub_cache) + candidates = [ + manifest_path(repo_type, repo_id, variant, hub_cache = hub_cache), + marker_path(repo_type, repo_id, variant, hub_cache = hub_cache), + ] + # Legacy unscoped state is shared: an unowned file belongs to the active + # cache (per _legacy_state_applies), so only purge it when it belongs to + # the cache being deleted -- else deleting an inactive cache would erase + # the active cache's resume/cancel state. + for path_factory in (manifest_path, marker_path): + legacy = path_factory(repo_type, repo_id, variant) + if legacy is not None and _legacy_state_applies(legacy, requested): + candidates.append(legacy) + paths = tuple(p for p in candidates if p is not None) + removed = False + for path in paths: + try: + if path.is_file(): + path.unlink() + removed = True + except OSError as exc: + logger.debug("Could not purge Hub state %s: %s", path, exc) + return removed -def purge_all_state_for_repo(repo_type: RepoType, repo_id: str) -> int: +def purge_all_state_for_repo( + repo_type: RepoType, + repo_id: str, + *, + hub_cache: Optional[str | Path] = None, +) -> int: """Remove the snapshot-level manifest + marker AND every variant-keyed manifest + marker for this repo. Used by the route delete handlers so scanner state never outlives the cache it described. Returns the count - of (repo, variant) triples that had any state on disk.""" + of (repo, variant) triples that had any state on disk. + + With ``hub_cache`` set, only that cache's scoped state (plus any legacy + unscoped file) is enumerated and removed, so deleting one cache's copy does + not clear another cache's resumable/cancel state.""" removed = 0 - if purge_state(repo_type, repo_id, None): + if purge_state(repo_type, repo_id, None, hub_cache = hub_cache): removed += 1 variants: set[str] = set() - for variant, _ in iter_variant_manifests(repo_type, repo_id): - variants.add(variant) - for variant, _ in iter_variant_markers(repo_type, repo_id): - variants.add(variant) + prefix = variant_filename_prefix(repo_type, repo_id) + if hub_cache is None: + search = [(p, True) for p in (manifests_dir(), cancelled_dir()) if p is not None] + else: + # This cache's scoped dir (parent of its scoped path) plus the legacy + # unscoped base; glob (not rglob) so other caches' dirs are not swept. + search = [] + for scoped, base in ( + (manifest_path(repo_type, repo_id, None, hub_cache = hub_cache), manifests_dir()), + (marker_path(repo_type, repo_id, None, hub_cache = hub_cache), cancelled_dir()), + ): + if scoped is not None: + search.append((scoped.parent, False)) + if base is not None: + search.append((base, False)) + for parent, recursive in search: + try: + entries = tuple( + parent.rglob(f"{prefix}*.json") if recursive else parent.glob(f"{prefix}*.json") + ) + except OSError: + continue + for entry in entries: + if not entry.is_file(): + continue + fallback = entry.stem[len(prefix) :] + variants.add(_variant_from_state_file(entry, fallback)) for variant in variants: - if purge_state(repo_type, repo_id, variant): + if purge_state(repo_type, repo_id, variant, hub_cache = hub_cache): removed += 1 return removed @@ -453,35 +666,83 @@ def _variant_from_state_file(path: Path, fallback: str) -> str: def _iter_variant_state_files( - parent: Optional[Path], repo_type: RepoType, repo_id: str + parent: Optional[Path], + repo_type: RepoType, + repo_id: str, + hub_cache: Optional[str | Path], + *, + cancel_markers: bool, ) -> Iterator[tuple[str, Path]]: if parent is None: return - prefix = variant_filename_prefix(repo_type, repo_id) - try: - entries = list(parent.iterdir()) - except OSError: + path_factory = marker_path if cancel_markers else manifest_path + requested = _canonical_hub_cache(hub_cache) + scoped_probe = path_factory( + repo_type, + repo_id, + None, + hub_cache = requested, + ) + if scoped_probe is None: return - for entry in entries: - if not entry.is_file() or not entry.name.endswith(".json"): + prefix = variant_filename_prefix(repo_type, repo_id) + seen: set[str] = set() + for directory, legacy in ((scoped_probe.parent, False), (parent, True)): + if legacy and directory == scoped_probe.parent: continue - stem = entry.name[: -len(".json")] - if not stem.lower().startswith(prefix): + try: + entries = list(directory.iterdir()) + except OSError: continue - variant = stem[len(prefix) :] - if variant: - yield _variant_from_state_file(entry, variant), entry + for entry in entries: + if not entry.is_file() or not entry.name.endswith(".json"): + continue + stem = entry.name[: -len(".json")] + if not stem.lower().startswith(prefix) or entry.name in seen: + continue + if legacy and not _legacy_state_applies( + entry, + requested, + fail_closed = cancel_markers, + ): + continue + fallback = stem[len(prefix) :] + if fallback: + seen.add(entry.name) + yield _variant_from_state_file(entry, fallback), entry -def iter_variant_manifests(repo_type: RepoType, repo_id: str) -> Iterator[tuple[str, Path]]: +def iter_variant_manifests( + repo_type: RepoType, + repo_id: str, + *, + hub_cache: Optional[str | Path] = None, +) -> Iterator[tuple[str, Path]]: """Yield (variant, manifest_path) for every variant-keyed manifest written for this repo. Used by is_gguf_repo_partial to enumerate all variants present on disk so the all-variants-broken gate can run.""" - yield from _iter_variant_state_files(manifests_dir(), repo_type, repo_id) + yield from _iter_variant_state_files( + manifests_dir(), + repo_type, + repo_id, + hub_cache, + cancel_markers = False, + ) -def iter_variant_markers(repo_type: RepoType, repo_id: str) -> Iterator[tuple[str, Path]]: +def iter_variant_markers( + repo_type: RepoType, + repo_id: str, + *, + hub_cache: Optional[str | Path] = None, +) -> Iterator[tuple[str, Path]]: """Yield (variant, marker_path) for every variant-keyed cancel marker. Companion to iter_variant_manifests: catches variants cancelled before download-start ever wrote a manifest (very early failures).""" - yield from _iter_variant_state_files(cancelled_dir(), repo_type, repo_id) + yield from _iter_variant_state_files( + cancelled_dir(), + repo_type, + repo_id, + hub_cache, + cancel_markers = True, + ) diff --git a/studio/backend/hub/utils/download_registry.py b/studio/backend/hub/utils/download_registry.py index b6bdee3bce..9e2b7d1a6d 100644 --- a/studio/backend/hub/utils/download_registry.py +++ b/studio/backend/hub/utils/download_registry.py @@ -129,6 +129,8 @@ def write_worker_breadcrumb(key: str, pid: int, metadata: Optional["DownloadMeta "cancel_marker_transport": metadata.cancel_marker_transport if metadata is not None else None, + "hub_cache": metadata.hub_cache if metadata is not None else None, + "xet_cache": metadata.xet_cache if metadata is not None else None, } tmp = path.with_name(f".{path.name}.tmp-{pid}") try: @@ -236,6 +238,7 @@ def _settle_orphaned_download( repo_id: Optional[str], variant: Optional[str], transport: Optional[str], + hub_cache: Optional[str] = None, ) -> None: """Persist a cancel marker for a reaped orphan still mid-download so the next launch settles it to a resumable "cancelled" state instead of a phantom-running @@ -251,18 +254,42 @@ def _settle_orphaned_download( return from hub.utils import download_manifest - manifest = download_manifest.read_manifest(repo_type, repo_id, variant) + cache_root = Path(hub_cache) if isinstance(hub_cache, str) and hub_cache else None + + manifest = download_manifest.read_manifest( + repo_type, + repo_id, + variant, + hub_cache = cache_root, + ) if repo_type == "model" and variant and manifest is None: return if manifest is None: - if not has_active_incomplete_blobs(repo_type, repo_id): + if not has_active_incomplete_blobs(repo_type, repo_id, root = cache_root): return else: - if _manifest_verifies_against_active_cache(repo_type, repo_id, manifest): + if _manifest_verifies_against_active_cache( + repo_type, + repo_id, + manifest, + root = cache_root, + ): return - if not _manifest_has_active_incomplete_blobs(repo_type, repo_id, manifest): + if not _manifest_has_active_incomplete_blobs( + repo_type, + repo_id, + manifest, + root = cache_root, + ): return - persist_cancel_marker(repo_type, repo_id, variant, transport, logger = logger) + persist_cancel_marker( + repo_type, + repo_id, + variant, + transport, + hub_cache = hub_cache, + logger = logger, + ) def reap_orphan_workers() -> None: @@ -309,6 +336,7 @@ def reap_orphan_workers() -> None: repo_id, data.get("variant"), data.get("cancel_marker_transport") or data.get("transport"), + data.get("hub_cache"), ) except Exception as exc: logger.debug("Reaper failed for breadcrumb %s: %s", entry, exc) @@ -355,8 +383,13 @@ def _purge_incomplete_blobs( return removed -def _iter_active_snapshot_dirs(repo_type: str, repo_id: str) -> Iterator[Path]: - for entry in iter_active_repo_cache_dirs(repo_type, repo_id): +def _iter_active_snapshot_dirs( + repo_type: str, + repo_id: str, + *, + root: Optional[Path] = None, +) -> Iterator[Path]: + for entry in iter_active_repo_cache_dirs(repo_type, repo_id, root = root): snapshots_dir = entry / "snapshots" if not snapshots_dir.is_dir(): continue @@ -369,24 +402,41 @@ def _iter_active_snapshot_dirs(repo_type: str, repo_id: str) -> Iterator[Path]: yield snapshot -def _manifest_verifies_against_active_cache(repo_type: str, repo_id: str, manifest) -> bool: +def _manifest_verifies_against_active_cache( + repo_type: str, + repo_id: str, + manifest, + *, + root: Optional[Path] = None, +) -> bool: from hub.utils import download_manifest - for snapshot_dir in _iter_active_snapshot_dirs(repo_type, repo_id): + for snapshot_dir in _iter_active_snapshot_dirs(repo_type, repo_id, root = root): if download_manifest.verify_against_disk(manifest, snapshot_dir).ok: return True return False -def _manifest_has_active_incomplete_blobs(repo_type: str, repo_id: str, manifest) -> bool: +def _manifest_has_active_incomplete_blobs( + repo_type: str, + repo_id: str, + manifest, + *, + root: Optional[Path] = None, +) -> bool: if not getattr(manifest, "variant", None): - return has_active_incomplete_blobs(repo_type, repo_id) + return has_active_incomplete_blobs(repo_type, repo_id, root = root) expected_hashes = frozenset( expected.sha256 for expected in manifest.expected_files if expected.sha256 ) if not expected_hashes: - return has_active_incomplete_blobs(repo_type, repo_id) + return has_active_incomplete_blobs(repo_type, repo_id, root = root) return bool( - incomplete_blob_hashes(repo_type, repo_id, active_only = True).intersection(expected_hashes) + incomplete_blob_hashes( + repo_type, + repo_id, + active_only = True, + root = root, + ).intersection(expected_hashes) ) @@ -459,6 +509,7 @@ def prepare_cache_for_transport( only_blob_hashes: Optional[frozenset[str]] = None, companion_blob_hashes: Optional[frozenset[str]] = None, protected_blob_hashes: Optional[frozenset[str]] = None, + root: Optional[Path] = None, ) -> int: """Guarantee any pre-existing ``.incomplete`` blobs are SAFE to resume under *mode*. Returns the number of partial blobs purged for untrusted provenance. @@ -485,14 +536,13 @@ def prepare_cache_for_transport( they are excluded from every purge so a shared companion is never deleted mid-write. - Scope: only the active ``HF_HUB_CACHE`` root is inspected. That suffices for - resume safety because ``snapshot_download`` runs without a ``cache_dir`` - override and so can only read or resume a ``.incomplete`` under this same - active root. Markers are written for the new mode before returning. + Scope: ``root`` selects the cache captured by the caller. It defaults to the + active ``HF_HUB_CACHE`` root for workers that inherit their cache through + the environment. Markers are written for the new mode before returning. """ if mode not in VALID_TRANSPORTS: raise ValueError(f"Invalid transport mode: {mode!r}") - root = hf_cache_root(create = True) + root = hf_cache_root(create = True) if root is None else hf_cache_root(create = True, root = root) if root is None: return 0 target = target_dir_name(repo_type, repo_id) @@ -618,10 +668,11 @@ def incomplete_blob_hashes( repo_id: str, *, active_only: bool = False, + root: Optional[Path] = None, ) -> set[str]: out: set[str] = set() entries = ( - iter_active_repo_cache_dirs(repo_type, repo_id) + iter_active_repo_cache_dirs(repo_type, repo_id, root = root) if active_only else iter_repo_cache_dirs(repo_type, repo_id) ) @@ -638,16 +689,24 @@ def incomplete_blob_hashes( return out -def completed_blob_bytes(repo_type: str, repo_id: str, blob_hashes: frozenset[str]) -> int: - """Sum finalized blob bytes for *blob_hashes* in the active HF cache root. +def completed_blob_bytes( + repo_type: str, + repo_id: str, + blob_hashes: frozenset[str], + *, + root: Optional[Path] = None, +) -> int: + """Sum finalized blob bytes for *blob_hashes* in a single HF cache root. - A worker only writes to the active ``HF_HUB_CACHE`` root, so a baseline must - ignore legacy/default roots that ``snapshot_download`` won't reuse this run. + A worker only writes to its captured ``HF_HUB_CACHE`` root, so a baseline + must be scoped to that root (``root``), not re-resolved to whatever cache is + active now; otherwise a runtime cache switch makes the retry baseline count + bytes from the wrong disk. """ if not blob_hashes: return 0 total = 0 - for entry in iter_active_repo_cache_dirs(repo_type, repo_id): + for entry in iter_active_repo_cache_dirs(repo_type, repo_id, root = root): blobs_dir = entry / "blobs" if not blobs_dir.is_dir(): continue @@ -712,6 +771,8 @@ class DownloadMetadata: # Bytes already complete before this job started; not counted as this run's # progress. completed_baseline_bytes: int = 0 + hub_cache: Optional[str] = None + xet_cache: Optional[str] = None @dataclass(frozen = True) @@ -752,6 +813,7 @@ def persist_cancel_marker( variant: Optional[str], transport: Optional[str], *, + hub_cache: Optional[str] = None, logger = logger, ) -> None: if not repo_type or not repo_id: @@ -763,6 +825,7 @@ def persist_cancel_marker( repo_id, variant, transport = transport, + hub_cache = hub_cache, ): logger.debug("write_cancel_marker returned False for %s", repo_id) except Exception as exc: @@ -971,6 +1034,7 @@ class DownloadRegistry: metadata_to_persist.repo_id, metadata_to_persist.variant, metadata_to_persist.transport, + hub_cache = metadata_to_persist.hub_cache, ) return False @@ -1033,6 +1097,8 @@ class DownloadRegistry: replace_active: bool = False, metadata_transport: Optional[str] = None, cancel_marker_transport: Optional[str] = None, + hub_cache: Optional[str] = None, + xet_cache: Optional[str] = None, ) -> tuple[bool, str]: key = normalize_job_key(key) repo = _repo_of_key(key) @@ -1106,6 +1172,8 @@ class DownloadRegistry: 0, int(completed_baseline_bytes or 0), ), + hub_cache = hub_cache, + xet_cache = xet_cache, ) if cancel_marker_transport is not None: self._cancel_marker_transports[key] = cancel_marker_transport @@ -1386,6 +1454,7 @@ class DownloadRegistry: metadata.repo_id, metadata.variant, metadata.cancel_marker_transport or metadata.transport, + hub_cache = metadata.hub_cache, ) reaped: list[tuple[str, subprocess.Popen, Optional[DownloadMetadata]]] = [] for key, proc, metadata in live: @@ -1401,6 +1470,7 @@ class DownloadRegistry: metadata.repo_id, metadata.variant, metadata.cancel_marker_transport or metadata.transport, + hub_cache = metadata.hub_cache, ) continue reaped.append((key, proc, metadata)) @@ -1421,6 +1491,7 @@ class DownloadRegistry: metadata.repo_id, metadata.variant, metadata.cancel_marker_transport or metadata.transport, + hub_cache = metadata.hub_cache, ) diff --git a/studio/backend/hub/utils/gguf.py b/studio/backend/hub/utils/gguf.py index 2e3de125f1..eb768db5d6 100644 --- a/studio/backend/hub/utils/gguf.py +++ b/studio/backend/hub/utils/gguf.py @@ -253,11 +253,16 @@ def _env_offline() -> bool: ) or os.environ.get("TRANSFORMERS_OFFLINE", "").lower() in ("1", "true", "yes") -def iter_hf_cache_snapshots(repo_id: str): - from hub.utils.hf_cache_state import iter_repo_cache_dirs +def iter_hf_cache_snapshots(repo_id: str, root: Optional[Path] = None): + from hub.utils.hf_cache_state import iter_active_repo_cache_dirs, iter_repo_cache_dirs snapshots: list[Path] = [] - for repo_dir in iter_repo_cache_dirs("model", repo_id): + repo_dirs = ( + iter_active_repo_cache_dirs("model", repo_id, root = root) + if root is not None + else iter_repo_cache_dirs("model", repo_id) + ) + for repo_dir in repo_dirs: snapshots_dir = repo_dir / "snapshots" if not snapshots_dir.is_dir(): continue @@ -276,12 +281,17 @@ def iter_hf_cache_snapshots(repo_id: str): yield from snapshots -def list_empty_gguf_variant_dirs(repo_id: str) -> set[str]: +def list_empty_gguf_variant_dirs(repo_id: str, root: Optional[Path] = None) -> set[str]: """Quant labels present only as an EMPTY snapshot ``/`` folder (an interrupted split download); a quant with shards in any snapshot is excluded.""" empty: dict[str, str] = {} nonempty: set[str] = set() - for snapshot in iter_hf_cache_snapshots(repo_id): + snapshots = ( + iter_hf_cache_snapshots(repo_id, root = root) + if root is not None + else iter_hf_cache_snapshots(repo_id) + ) + for snapshot in snapshots: try: entries = list(snapshot.iterdir()) except OSError: @@ -303,8 +313,15 @@ def list_empty_gguf_variant_dirs(repo_id: str) -> set[str]: return {label for key, label in empty.items() if key not in nonempty} -def list_gguf_variants_from_hf_cache(repo_id: str) -> Optional[tuple[list[GgufVariantInfo], bool]]: - for snapshot in iter_hf_cache_snapshots(repo_id): +def list_gguf_variants_from_hf_cache( + repo_id: str, root: Optional[Path] = None +) -> Optional[tuple[list[GgufVariantInfo], bool]]: + snapshots = ( + iter_hf_cache_snapshots(repo_id, root = root) + if root is not None + else iter_hf_cache_snapshots(repo_id) + ) + for snapshot in snapshots: variants, has_vision = list_local_gguf_variants(str(snapshot)) if variants or has_vision: return variants, has_vision @@ -312,7 +329,7 @@ def list_gguf_variants_from_hf_cache(repo_id: str) -> Optional[tuple[list[GgufVa def list_partial_gguf_variants_from_state( - repo_id: str, + repo_id: str, hub_cache: Optional[Path] = None ) -> Optional[tuple[list[GgufVariantInfo], bool]]: """Reconstruct GGUF variants from download manifests/markers alone. @@ -328,10 +345,26 @@ def list_partial_gguf_variants_from_state( # original-casing label over a lowercased cancel marker for the same variant. seen: set[str] = set() ordered: list[str] = [] - for source in ( - download_manifest.iter_variant_manifests("model", repo_id), - download_manifest.iter_variant_markers("model", repo_id), - ): + sources = ( + ( + download_manifest.iter_variant_manifests("model", repo_id), + download_manifest.iter_variant_markers("model", repo_id), + ) + if hub_cache is None + else ( + download_manifest.iter_variant_manifests( + "model", + repo_id, + hub_cache = hub_cache, + ), + download_manifest.iter_variant_markers( + "model", + repo_id, + hub_cache = hub_cache, + ), + ) + ) + for source in sources: for variant, _path in source: key = variant.lower() if key not in seen: @@ -343,7 +376,16 @@ def list_partial_gguf_variants_from_state( variants: list[GgufVariantInfo] = [] has_vision = False for variant in ordered: - manifest = download_manifest.read_manifest("model", repo_id, variant) + manifest = ( + download_manifest.read_manifest("model", repo_id, variant) + if hub_cache is None + else download_manifest.read_manifest( + "model", + repo_id, + variant, + hub_cache = hub_cache, + ) + ) main_filename: Optional[str] = None size_bytes = 0 companion_bytes = 0 diff --git a/studio/backend/hub/utils/hf_cache_state.py b/studio/backend/hub/utils/hf_cache_state.py index 22c948b683..49a28c813c 100644 --- a/studio/backend/hub/utils/hf_cache_state.py +++ b/studio/backend/hub/utils/hf_cache_state.py @@ -29,12 +29,10 @@ def _safe_is_dir(path: Path) -> bool: return False -def hf_cache_root(*, create: bool = False) -> Optional[Path]: - try: - from huggingface_hub import constants as hf_constants - except ImportError: - return None - root = Path(hf_constants.HF_HUB_CACHE) +def hf_cache_root(*, create: bool = False, root: Optional[Path] = None) -> Optional[Path]: + from utils.hf_cache_settings import get_hf_cache_paths + + root = root or get_hf_cache_paths().hub_cache if create: try: root.mkdir(parents = True, exist_ok = True) @@ -46,6 +44,7 @@ def hf_cache_root(*, create: bool = False) -> Optional[Path]: def hf_cache_roots() -> list[Path]: from hub.utils.paths import hf_default_cache_dir, legacy_hf_cache_dir + from utils.hf_cache_settings import known_hf_hub_caches roots: list[Path] = [] seen: set[str] = set() @@ -62,7 +61,8 @@ def hf_cache_roots() -> list[Path]: seen.add(key) roots.append(path) - _add(hf_cache_root()) + for configured in known_hf_hub_caches(): + _add(configured) _add(legacy_hf_cache_dir()) _add(hf_default_cache_dir()) return roots @@ -181,12 +181,22 @@ def iter_repo_cache_dirs(repo_type: str, repo_id: str) -> Iterator[Path]: continue -def iter_destructive_repo_cache_dirs(repo_type: str, repo_id: str) -> Iterator[Path]: +def iter_destructive_repo_cache_dirs( + repo_type: str, + repo_id: str, + *, + root: Optional[Path] = None, +) -> Iterator[Path]: target = repo_cache_dir_name(repo_type, repo_id) folded_target = target.lower() - for root in hf_cache_roots(): + if root is not None: + scoped = hf_cache_root(root = root) + bases = [scoped] if scoped is not None else [] + else: + bases = hf_cache_roots() + for base in bases: try: - entries = [entry for entry in root.iterdir() if entry.name.lower() == folded_target] + entries = [entry for entry in base.iterdir() if entry.name.lower() == folded_target] except OSError: continue matched_names = resolve_destructive_case_matches( @@ -200,8 +210,13 @@ def iter_destructive_repo_cache_dirs(repo_type: str, repo_id: str) -> Iterator[P yield entry -def iter_active_repo_cache_dirs(repo_type: str, repo_id: str) -> Iterator[Path]: - root = hf_cache_root() +def iter_active_repo_cache_dirs( + repo_type: str, + repo_id: str, + *, + root: Optional[Path] = None, +) -> Iterator[Path]: + root = hf_cache_root(root = root) if root is None: return target = target_dir_name(repo_type, repo_id) @@ -218,12 +233,13 @@ def preferred_repo_cache_dirs( repo_id: str, *, force_active: bool = False, + active_root: Optional[Path] = None, ) -> list[Path]: - active_entries = list(iter_active_repo_cache_dirs(repo_type, repo_id)) + active_entries = list(iter_active_repo_cache_dirs(repo_type, repo_id, root = active_root)) if active_entries: return active_entries if force_active: - root = hf_cache_root() + root = hf_cache_root(root = active_root) if root is not None: canonical = repo_cache_dir_name(repo_type, repo_id) return [root / canonical] @@ -237,8 +253,13 @@ def has_incomplete_blobs(repo_type: str, repo_id: str) -> bool: return False -def has_active_incomplete_blobs(repo_type: str, repo_id: str) -> bool: - for entry in iter_active_repo_cache_dirs(repo_type, repo_id): +def has_active_incomplete_blobs( + repo_type: str, + repo_id: str, + *, + root: Optional[Path] = None, +) -> bool: + for entry in iter_active_repo_cache_dirs(repo_type, repo_id, root = root): if repo_cache_dir_has_incomplete_blobs(entry): return True return False @@ -273,9 +294,14 @@ def _prune_empty_dirs(root: Path) -> bool: return removed -def purge_partial_repo(repo_type: str, repo_id: str) -> bool: +def purge_partial_repo( + repo_type: str, + repo_id: str, + *, + root: Optional[Path] = None, +) -> bool: removed = False - for entry in iter_destructive_repo_cache_dirs(repo_type, repo_id): + for entry in iter_destructive_repo_cache_dirs(repo_type, repo_id, root = root): blobs_dir = entry / "blobs" if blobs_dir.is_dir(): for blob in blobs_dir.iterdir(): @@ -290,9 +316,14 @@ def purge_partial_repo(repo_type: str, repo_id: str) -> bool: return removed -def purge_repo_cache_dirs(repo_type: str, repo_id: str) -> bool: +def purge_repo_cache_dirs( + repo_type: str, + repo_id: str, + *, + root: Optional[Path] = None, +) -> bool: removed = False - for entry in iter_destructive_repo_cache_dirs(repo_type, repo_id): + for entry in iter_destructive_repo_cache_dirs(repo_type, repo_id, root = root): try: if entry.is_symlink() or not entry.is_dir(): continue @@ -301,3 +332,59 @@ def purge_repo_cache_dirs(repo_type: str, repo_id: str) -> bool: except FileNotFoundError: continue return removed + + +def scoped_delete_root(repo_type: str, repo_id: str, cache_path: Optional[str]) -> Optional[Path]: + """Resolve the single cache root a delete of this repo may touch. + + Returns the active hub cache when *cache_path* is falsy, the owning cache + root when *cache_path* points inside a known cache, or ``None`` when + *cache_path* is set but not inside any known cache (caller should reject). + This keeps a delete of one inventory row from removing copies in other, + previously selected caches. + """ + from utils.hf_cache_settings import get_hf_cache_paths + + if not cache_path: + return Path(get_hf_cache_paths().hub_cache).resolve(strict = False) + try: + resolved = Path(cache_path).expanduser().resolve(strict = False) + except (OSError, RuntimeError, ValueError): + return None + expected = repo_cache_dir_name(repo_type, repo_id).lower() + repo_dir = next( + ( + candidate + for candidate in (resolved, *resolved.parents) + if candidate.name.lower() == expected + ), + None, + ) + if repo_dir is None: + return None + allowed = {r.resolve(strict = False) for r in hf_cache_roots()} + root = repo_dir.parent.resolve(strict = False) + return root if root in allowed else None + + +def resolve_delete_target_root( + repo_type: str, repo_id: str, cache_path: Optional[str], owner_roots +) -> Optional[Path]: + """Pick the single cache root a delete of this repo should target. + + An explicit *cache_path* wins (``None`` when it is not a known cache, so the + caller can reject it). Otherwise prefer the active cache when it holds a + copy, else the sole cache that does -- so a model that lives only in a + previously selected cache stays deletable while other caches are untouched. + """ + if cache_path: + return scoped_delete_root(repo_type, repo_id, cache_path) + from utils.hf_cache_settings import get_hf_cache_paths + + active = Path(get_hf_cache_paths().hub_cache).resolve(strict = False) + roots = list(owner_roots) + if active in roots: + return active + if len(roots) == 1: + return roots[0] + return active diff --git a/studio/backend/hub/utils/inventory_scan.py b/studio/backend/hub/utils/inventory_scan.py index 57ad7f6655..058fdf9b65 100644 --- a/studio/backend/hub/utils/inventory_scan.py +++ b/studio/backend/hub/utils/inventory_scan.py @@ -36,7 +36,7 @@ from hub.utils.state_dir import RepoType from hub.utils.hf_cache_state import ( INCOMPLETE_SUFFIX, has_incomplete_blobs, - hf_cache_root, + hf_cache_roots, iter_repo_cache_dirs, latest_snapshot_dir, repo_cache_dir_has_incomplete_blobs, @@ -127,33 +127,13 @@ def all_hf_cache_scans() -> list: def _compute_all_hf_cache_scans() -> list: from huggingface_hub import scan_cache_dir - from hub.utils.paths import legacy_hf_cache_dir, hf_default_cache_dir scans: list = [] - seen: set[str] = set() - try: - from huggingface_hub.constants import HF_HUB_CACHE - - active = Path(HF_HUB_CACHE).resolve() - seen.add(str(active)) - if active.is_dir(): - scans.append(scan_cache_dir()) - except Exception as exc: - logger.warning("Could not scan active HF cache: %s", exc) - - for extra_fn in (legacy_hf_cache_dir, hf_default_cache_dir): + for cache_root in hf_cache_roots(): try: - extra = extra_fn() - # is_dir()/resolve() can raise on an inaccessible path; skip it. - if not extra.is_dir(): - continue - resolved = str(extra.resolve()) - if resolved in seen: - continue - seen.add(resolved) - scans.append(scan_cache_dir(cache_dir = str(extra))) + scans.append(scan_cache_dir(cache_dir = str(cache_root))) except Exception as exc: - logger.warning("Could not scan HF cache %s: %s", extra_fn.__name__, exc) + logger.warning("Could not scan HF cache %s: %s", cache_root, exc) return scans @@ -224,16 +204,8 @@ def _compose_partial(*signals: Callable[[], bool]) -> bool: return any(signal() for signal in signals) -def _state_applies_to_repo_cache_dir(repo_cache_dir: Optional[Path]) -> bool: - if repo_cache_dir is None: - return True - root = hf_cache_root() - if root is None: - return False - try: - return repo_cache_dir.resolve().parent == root.resolve() - except OSError: - return False +def _hub_cache_for_repo_dir(repo_cache_dir: Optional[Path]) -> Optional[Path]: + return repo_cache_dir.parent if repo_cache_dir is not None else None def _legacy_partial( @@ -285,12 +257,24 @@ def _repo_cache_dir_has_non_gguf_broken_snapshot_symlinks(repo_cache_dir: Path) return False -def _gguf_variant_manifest_blob_hashes(repo_id: str) -> frozenset[str]: +def _gguf_variant_manifest_blob_hashes( + repo_id: str, repo_cache_dir: Optional[Path] = None +) -> frozenset[str]: from hub.utils import download_manifest hashes: set[str] = set() - for variant, _path in download_manifest.iter_variant_manifests("model", repo_id): - manifest = download_manifest.read_manifest("model", repo_id, variant) + hub_cache = _hub_cache_for_repo_dir(repo_cache_dir) + for variant, _path in download_manifest.iter_variant_manifests( + "model", + repo_id, + hub_cache = hub_cache, + ): + manifest = download_manifest.read_manifest( + "model", + repo_id, + variant, + hub_cache = hub_cache, + ) if manifest is None: continue for expected in manifest.expected_files: @@ -315,7 +299,7 @@ def _snapshot_legacy_partial( ) -> bool: if repo_type != "model": return _legacy_partial(repo_type, repo_id, repo_cache_dir) - ignored_hashes = _gguf_variant_manifest_blob_hashes(repo_id) + ignored_hashes = _gguf_variant_manifest_blob_hashes(repo_id, repo_cache_dir) if repo_cache_dir is not None: return _repo_cache_dir_has_snapshot_legacy_partial( repo_cache_dir, @@ -375,9 +359,12 @@ def _manifest_partial( ) -> bool: from hub.utils import download_manifest - if not _state_applies_to_repo_cache_dir(repo_cache_dir): - return False - manifest = download_manifest.read_manifest(repo_type, repo_id, variant) + manifest = download_manifest.read_manifest( + repo_type, + repo_id, + variant, + hub_cache = _hub_cache_for_repo_dir(repo_cache_dir), + ) if manifest is None: return False resolved = ( @@ -452,10 +439,13 @@ def is_snapshot_partial( A manifest without a resolvable snapshot is partial: the worker got far enough to record expectations but did not leave a usable snapshot.""" from hub.utils import download_manifest - - state_applies = _state_applies_to_repo_cache_dir(repo_cache_dir) return _compose_partial( - lambda: state_applies and download_manifest.has_cancel_marker(repo_type, repo_id, None), + lambda: download_manifest.has_cancel_marker( + repo_type, + repo_id, + None, + hub_cache = _hub_cache_for_repo_dir(repo_cache_dir), + ), lambda: _snapshot_legacy_partial(repo_type, repo_id, repo_cache_dir), lambda: _manifest_partial( repo_type, @@ -484,10 +474,13 @@ def is_variant_partial( caller is checking many variants of the same repo (see is_gguf_repo_partial for that usage).""" from hub.utils import download_manifest - - state_applies = _state_applies_to_repo_cache_dir(repo_cache_dir) return _compose_partial( - lambda: state_applies and download_manifest.has_cancel_marker("model", repo_id, variant), + lambda: download_manifest.has_cancel_marker( + "model", + repo_id, + variant, + hub_cache = _hub_cache_for_repo_dir(repo_cache_dir), + ), lambda: bool( incomplete_blob_hashes and variant_blob_hashes @@ -526,22 +519,38 @@ def is_gguf_repo_partial(repo_id: str, repo_cache_dir: Optional[Path] = None) -> from hub.utils import download_manifest has_legacy_partial = _legacy_partial("model", repo_id, repo_cache_dir) - state_applies = _state_applies_to_repo_cache_dir(repo_cache_dir) snapshot_dir = resolve_snapshot_dir_for_scan( "model", repo_id, repo_cache_dir, ) variants: set[str] = set(_completed_gguf_variants(snapshot_dir)) - if state_applies: - for variant, _path in download_manifest.iter_variant_manifests( - "model", - repo_id, + hub_cache = _hub_cache_for_repo_dir(repo_cache_dir) + for variant, _path in download_manifest.iter_variant_manifests( + "model", + repo_id, + hub_cache = hub_cache, + ): + if ( + download_manifest.read_manifest( + "model", + repo_id, + variant, + hub_cache = hub_cache, + ) + is not None ): variants.add(variant) - for variant, _path in download_manifest.iter_variant_markers( + for variant, _path in download_manifest.iter_variant_markers( + "model", + repo_id, + hub_cache = hub_cache, + ): + if download_manifest.has_cancel_marker( "model", repo_id, + variant, + hub_cache = hub_cache, ): variants.add(variant) if not variants: @@ -576,14 +585,19 @@ def partial_transport_for( available.""" from hub.utils import download_manifest - if not _state_applies_to_repo_cache_dir(repo_cache_dir): - return None + hub_cache = _hub_cache_for_repo_dir(repo_cache_dir) marker_transport = download_manifest.read_cancel_marker_transport( repo_type, repo_id, variant, + hub_cache = hub_cache, ) if marker_transport is not None: return marker_transport - manifest = download_manifest.read_manifest(repo_type, repo_id, variant) + manifest = download_manifest.read_manifest( + repo_type, + repo_id, + variant, + hub_cache = hub_cache, + ) return manifest.transport if manifest is not None else None diff --git a/studio/backend/hub/utils/paths.py b/studio/backend/hub/utils/paths.py index 5435202565..81621edcf9 100644 --- a/studio/backend/hub/utils/paths.py +++ b/studio/backend/hub/utils/paths.py @@ -277,12 +277,8 @@ def _memo_drop(memo_key: tuple[str, str]) -> None: def _hf_hub_cache_dir() -> Path: - try: - from huggingface_hub.constants import HF_HUB_CACHE - return Path(HF_HUB_CACHE) - except Exception as exc: - logger.debug("Could not read huggingface_hub HF_HUB_CACHE, using default: %s", exc) - return Path.home() / ".cache" / "huggingface" / "hub" + from utils.hf_cache_settings import get_hf_cache_paths + return get_hf_cache_paths().hub_cache def _hf_hub_cache_dirs() -> list[Path]: @@ -300,7 +296,10 @@ def _hf_hub_cache_dirs() -> list[Path]: seen.add(key) roots.append(resolved) - _add(_hf_hub_cache_dir()) + from utils.hf_cache_settings import known_hf_hub_caches + + for configured in known_hf_hub_caches(): + _add(configured) try: _add(legacy_hf_cache_dir()) _add(hf_default_cache_dir()) diff --git a/studio/backend/hub/utils/state_dir.py b/studio/backend/hub/utils/state_dir.py index 898c03c87d..4650b97381 100644 --- a/studio/backend/hub/utils/state_dir.py +++ b/studio/backend/hub/utils/state_dir.py @@ -8,9 +8,10 @@ so it survives ``huggingface-cli delete-cache`` and any other HF-side cache lifecycle. Two subdirectories: /hub-state/ - manifests/ .json per-download expected-files manifest - cancelled/ .json per-download cancel marker + manifests/cache-/.json expected-files manifest + cancelled/cache-/.json cancel marker +The cache digest isolates state for the same repo across selectable Hub caches. The ```` mirrors HF's cache dir naming while the resulting manifest, cancel-marker, and atomic-write temp filenames fit common filesystem basename limits. Very long repo IDs use a stable hash in the state key: @@ -29,6 +30,7 @@ configuration failure. from __future__ import annotations import hashlib +import os import re from pathlib import Path from typing import Literal, Optional, get_args @@ -55,6 +57,7 @@ _STATE_EXTENSION = ".json" # _atomic_write_json writes "..tmp-<8hex>" beside the final file. _ATOMIC_WRITE_TMP_OVERHEAD = len(".") + len(".tmp-") + 8 _MAX_VARIANT_FRAGMENT_LENGTH = 64 +_CACHE_SCOPE_DIGEST_LENGTH = 32 def state_root() -> Optional[Path]: @@ -130,13 +133,32 @@ def _entry_key(repo_type: RepoType, repo_id: str, variant: Optional[str]) -> str return f"{variant_filename_prefix(repo_type, repo_id)}{variant_fragment}" +def _cache_scope(parent: Path, hub_cache: Optional[str | Path]) -> Optional[Path]: + if hub_cache is None: + return parent + normalized = os.path.normcase(str(Path(hub_cache).expanduser())) + digest = hashlib.sha256(normalized.encode("utf-8")).hexdigest()[:_CACHE_SCOPE_DIGEST_LENGTH] + scoped = parent / f"cache-{digest}" + try: + scoped.mkdir(parents = True, exist_ok = True) + except OSError as exc: + logger.debug("Could not create cache-scoped Hub state dir %s: %s", scoped, exc) + return None + return scoped + + def manifest_path( repo_type: RepoType, repo_id: str, variant: Optional[str] = None, + *, + hub_cache: Optional[str | Path] = None, ) -> Optional[Path]: """Path to the manifest file for this triple. May or may not exist.""" parent = _subdir(_MANIFESTS_SUBDIR) + if parent is None: + return None + parent = _cache_scope(parent, hub_cache) if parent is None: return None return parent / f"{_entry_key(repo_type, repo_id, variant)}.json" @@ -146,9 +168,14 @@ def marker_path( repo_type: RepoType, repo_id: str, variant: Optional[str] = None, + *, + hub_cache: Optional[str | Path] = None, ) -> Optional[Path]: """Path to the cancel-marker file for this triple. May or may not exist.""" parent = _subdir(_CANCELLED_SUBDIR) + if parent is None: + return None + parent = _cache_scope(parent, hub_cache) if parent is None: return None return parent / f"{_entry_key(repo_type, repo_id, variant)}.json" diff --git a/studio/backend/hub/workers/hf_download.py b/studio/backend/hub/workers/hf_download.py index e45357d311..9ff394b009 100644 --- a/studio/backend/hub/workers/hf_download.py +++ b/studio/backend/hub/workers/hf_download.py @@ -661,6 +661,7 @@ def _download_gguf_variant(repo_id: str, variant: str, hf_token: str | None, mod variant, plan.main_hashes, hf_token, + hub_cache = Path(snapshot_path).parents[2], ) except Exception as e: print( diff --git a/studio/backend/models/models.py b/studio/backend/models/models.py index 54e88fed58..df6725c9c9 100644 --- a/studio/backend/models/models.py +++ b/studio/backend/models/models.py @@ -178,6 +178,14 @@ class LocalModelInfo(BaseModel): None, description = "HF repo id for cached models, e.g. org/model", ) + active_cache: Optional[bool] = Field( + None, + description = "Whether an HF model belongs to the current download cache.", + ) + partial: bool = Field( + False, + description = "Whether the cached model has an incomplete download.", + ) model_format: Optional[str] = Field( None, description = "Detected weights format ('gguf' when known). Lets the UI " diff --git a/studio/backend/picker/service.py b/studio/backend/picker/service.py index 13065b2920..ccf9c3e152 100644 --- a/studio/backend/picker/service.py +++ b/studio/backend/picker/service.py @@ -22,6 +22,7 @@ from utils.models.model_config import ( _is_mmproj, _is_mtp_drafter, ) +from utils.hf_cache_settings import active_hf_hub_cache from utils.paths.path_utils import ( is_local_path, normalize_path, @@ -378,7 +379,12 @@ def read_default_chat_template( if _remote_exceeds_cap(rel): return None try: - path = hf_hub_download(resolved, rel, token = hf_token) + path = hf_hub_download( + resolved, + rel, + token = hf_token, + cache_dir = active_hf_hub_cache(), + ) return _read_bounded_text(Path(path), MAX_TEMPLATE_METADATA_BYTES) except Exception: return None diff --git a/studio/backend/routes/datasets.py b/studio/backend/routes/datasets.py index 5456080f34..331eb5e1e0 100644 --- a/studio/backend/routes/datasets.py +++ b/studio/backend/routes/datasets.py @@ -23,40 +23,6 @@ def _is_valid_repo_id(repo_id: str) -> bool: return bool(_VALID_REPO_ID.fullmatch(repo_id)) -_dataset_size_cache: dict[str, int] = {} - - -def _get_dataset_size_cached(repo_id: str) -> int: - if repo_id in _dataset_size_cache: - return _dataset_size_cache[repo_id] - try: - from huggingface_hub import dataset_info as hf_dataset_info - - info = hf_dataset_info(repo_id, token = None, files_metadata = True) - total = sum(s.size for s in info.siblings if getattr(s, "size", None)) - _dataset_size_cache[repo_id] = total - return total - except Exception: - return 0 - - -def _resolve_hf_cache_realpath(repo_dir: Path) -> Optional[str]: - """Resolved realpath for a HF cache repo dir: most-recent snapshot, else cache root. - - Mirrors routes/models.py; duplicated here to keep this module self-contained. - """ - try: - snapshots_dir = repo_dir / "snapshots" - if snapshots_dir.is_dir(): - snaps = [s for s in snapshots_dir.iterdir() if s.is_dir()] - if snaps: - latest = max(snaps, key = lambda s: s.stat().st_mtime) - return str(latest.resolve()) - return str(repo_dir.resolve()) - except Exception: - return None - - backend_path = Path(__file__).parent.parent.parent if str(backend_path) not in sys.path: sys.path.insert(0, str(backend_path)) @@ -64,6 +30,7 @@ if str(backend_path) not in sys.path: from utils.datasets import check_dataset_format from utils.upload_limits import get_upload_limit_bytes, get_upload_limit_label from auth.authentication import get_current_subject +from hub.dependencies import get_hf_token router = APIRouter() logger = get_logger(__name__) @@ -292,11 +259,13 @@ def _download_hf_metadata(*, repo_id: str, repo_files: list[str], token: str | N try: from huggingface_hub import hf_hub_download + from utils.hf_cache_settings import active_hf_hub_cache local_path = hf_hub_download( repo_id = repo_id, filename = metadata_file, repo_type = "dataset", token = token, + cache_dir = active_hf_hub_cache(), ) except Exception as exc: logger.warning(f"Could not read HF dataset metadata for {repo_id}: {exc}") @@ -525,77 +494,15 @@ def list_local_datasets( @router.get("/download-progress") async def get_dataset_download_progress( repo_id: str = Query(..., description = "HuggingFace dataset repo ID, e.g. 'unsloth/LaTeX_OCR'"), + hf_token: Optional[str] = Depends(get_hf_token), current_subject: str = Depends(get_current_subject), ): - """Return download progress for a HuggingFace dataset repo. - - Mirrors ``GET /api/models/download-progress`` but scans the - ``datasets--owner--name`` cache dir under HF_HUB_CACHE, where in-progress - download bytes are visible. Returns ``cache_path`` so the UI can show it. - """ - _empty = { - "downloaded_bytes": 0, - "expected_bytes": 0, - "progress": 0, - "cache_path": None, - } - try: - if not _is_valid_repo_id(repo_id): - return _empty - - from huggingface_hub import constants as hf_constants - - cache_dir = Path(hf_constants.HF_HUB_CACHE) - target = f"datasets--{repo_id.replace('/', '--')}".lower() - completed_bytes = 0 - in_progress_bytes = 0 - cache_path: Optional[str] = None - - if cache_dir.is_dir(): - for entry in cache_dir.iterdir(): - if entry.name.lower() != target: - continue - cache_path = _resolve_hf_cache_realpath(entry) - blobs_dir = entry / "blobs" - if not blobs_dir.is_dir(): - break - for f in blobs_dir.iterdir(): - if not f.is_file(): - continue - if f.name.endswith(".incomplete"): - in_progress_bytes += f.stat().st_size - else: - completed_bytes += f.stat().st_size - break - - downloaded_bytes = completed_bytes + in_progress_bytes - if downloaded_bytes == 0: - return {**_empty, "cache_path": cache_path} - - expected_bytes = _get_dataset_size_cached(repo_id) - if expected_bytes <= 0: - return { - "downloaded_bytes": downloaded_bytes, - "expected_bytes": 0, - "progress": 0, - "cache_path": cache_path, - } - - # 95% threshold (as in the model endpoint): HF blob dedup makes - # completed_bytes drift under expected_bytes; inter-file gaps look "done". - if completed_bytes >= expected_bytes * 0.95: - progress = 1.0 - else: - progress = min(downloaded_bytes / expected_bytes, 0.99) - return { - "downloaded_bytes": downloaded_bytes, - "expected_bytes": expected_bytes, - "progress": round(progress, 3), - "cache_path": cache_path, - } - except Exception as e: - logger.warning(f"Error checking dataset download progress for {repo_id}: {e}") - return _empty + """Compatibility route backed by the shared multi-cache progress service.""" + from hub.services.datasets import downloads + return await downloads.get_dataset_download_progress_response( + repo_id, + hf_token = hf_token, + ) @router.post("/check-format", response_model = CheckFormatResponse) diff --git a/studio/backend/routes/models.py b/studio/backend/routes/models.py index a5ce1a72f0..c225483acf 100644 --- a/studio/backend/routes/models.py +++ b/studio/backend/routes/models.py @@ -224,11 +224,8 @@ def derive_model_type( def _resolve_hf_cache_dir() -> Path: """Resolve local HF cache root used by hub downloads.""" - try: - from huggingface_hub.constants import HF_HUB_CACHE - return Path(HF_HUB_CACHE) - except Exception: - return Path.home() / ".cache" / "huggingface" / "hub" + from utils.hf_cache_settings import get_hf_cache_paths + return get_hf_cache_paths().hub_cache def _is_model_directory(d: Path) -> bool: @@ -370,10 +367,12 @@ def _scan_models_dir(models_dir: Path, *, limit: int | None = None) -> List[Loca return found -def _scan_hf_cache(cache_dir: Path) -> List[LocalModelInfo]: +def _scan_hf_cache(cache_dir: Path, *, active_cache: bool = True) -> List[LocalModelInfo]: if not cache_dir.exists() or not cache_dir.is_dir(): return [] + from hub.utils import inventory_scan as hf_cache_scan + found: List[LocalModelInfo] = [] for repo_dir in cache_dir.glob("models--*"): if not repo_dir.is_dir(): @@ -389,13 +388,21 @@ def _scan_hf_cache(cache_dir: Path) -> List[LocalModelInfo]: except OSError: updated_at = None + partial = hf_cache_scan.is_snapshot_partial("model", model_id, repo_dir) + partial = partial or hf_cache_scan.is_gguf_repo_partial(model_id, repo_dir) + + load_id = model_id + if not active_cache: + load_id = _resolve_hf_cache_realpath(repo_dir) or str(repo_dir.resolve()) found.append( LocalModelInfo( - id = model_id, + id = load_id, model_id = model_id, display_name = model_id.split("/")[-1], - path = str(repo_dir), + path = load_id if not active_cache else str(repo_dir), source = "hf_cache", + active_cache = active_cache, + partial = partial, updated_at = updated_at, ), ) @@ -776,26 +783,34 @@ def collect_local_models(models_root: Path) -> List[LocalModelInfo]: legacy_hf_cache_dir, lmstudio_model_dirs, ) + from utils.hf_cache_settings import known_hf_hub_caches hf_cache_dir = _resolve_hf_cache_dir() legacy_hf = legacy_hf_cache_dir() hf_default = hf_default_cache_dir() lm_dirs = lmstudio_model_dirs() - local_models = _scan_models_dir(models_root) + _scan_hf_cache(hf_cache_dir) - - # Resolve once; an inaccessible aux cache must skip that scan, not 500. - hf_cache_real = _safe_resolve(hf_cache_dir) - legacy_real = _safe_resolve(legacy_hf) - default_real = _safe_resolve(hf_default) - - # Scan legacy Unsloth HF cache for backward compatibility. - if _safe_is_dir(legacy_hf) and legacy_real != hf_cache_real: - local_models += _scan_hf_cache(legacy_hf) - - # Scan HF system default cache (may differ under env overrides). - if _safe_is_dir(hf_default) and default_real != hf_cache_real and default_real != legacy_real: - local_models += _scan_hf_cache(hf_default) + local_models = _scan_models_dir(models_root) + active_cache_real = _safe_resolve(hf_cache_dir) + active_cache_key = os.path.normcase(active_cache_real) if active_cache_real else None + seen_hf: set[str] = set() + for cache_dir in ( + hf_cache_dir, + *known_hf_hub_caches(), + legacy_hf, + hf_default, + ): + cache_real = _safe_resolve(cache_dir) + if cache_real is None: + continue + cache_key = os.path.normcase(str(cache_real)) + if cache_key in seen_hf: + continue + seen_hf.add(cache_key) + local_models += _scan_hf_cache( + cache_dir, + active_cache = cache_key == active_cache_key, + ) # Scan LM Studio directories. for lm_dir in lm_dirs: @@ -817,7 +832,7 @@ def collect_local_models(models_root: Path) -> List[LocalModelInfo]: m for m in ( _scan_models_dir(folder_path, limit = _MAX_MODELS_PER_FOLDER) - + _scan_hf_cache(folder_path) + + _scan_hf_cache(folder_path, active_cache = False) + _scan_lmstudio_dir(folder_path) ) if not any(p in (".studio_links", "ollama_links") for p in Path(m.path).parts) @@ -838,8 +853,18 @@ def collect_local_models(models_root: Path) -> List[LocalModelInfo]: # even when the model is also in the HF cache. deduped: dict[str, LocalModelInfo] = {} for model in local_models: - key = f"{model.id}\x00custom" if model.source == "custom" else model.id - if key not in deduped: + semantic_id = model.model_id if model.source == "hf_cache" and model.model_id else model.id + key = f"{semantic_id}\x00custom" if model.source == "custom" else semantic_id + existing = deduped.get(key) + prefer_model = existing is None + if existing is not None and model.source == existing.source == "hf_cache": + if model.partial != existing.partial: + prefer_model = not model.partial + elif bool(model.active_cache) != bool(existing.active_cache): + prefer_model = bool(model.active_cache) + else: + prefer_model = (model.updated_at or 0) > (existing.updated_at or 0) + if prefer_model: deduped[key] = model models = sorted( @@ -1202,10 +1227,7 @@ def _build_browse_allowlist( legacy_hf_cache_dir, well_known_model_dirs, ) - from utils.paths.external_media import ( - linux_run_media_mount_roots, - windows_drive_roots, - ) + from utils.paths import external_media from storage.studio_db import list_scan_folders candidates: list[Path] = [] @@ -1222,9 +1244,12 @@ def _build_browse_allowlist( _add(Path.home()) if media_roots is None: - media_roots = linux_run_media_mount_roots() + media_roots = [ + *external_media.linux_run_media_mount_roots(), + *external_media.macos_volume_roots(), + ] if drive_roots is None: - drive_roots = windows_drive_roots() + drive_roots = external_media.windows_drive_roots() for p in media_roots: _add(p) for p in drive_roots: @@ -1502,10 +1527,7 @@ def browse_folders( then hidden (if ``show_hidden=true``). """ from utils.paths import hf_default_cache_dir, well_known_model_dirs - from utils.paths.external_media import ( - linux_run_media_mount_roots, - windows_drive_roots, - ) + from utils.paths import external_media from storage.studio_db import ( contains_sensitive_path_component, is_denied_system_path, @@ -1514,8 +1536,11 @@ def browse_folders( # Probe removable-media and Windows drive roots once; the allowlist and # chips reuse the result so a disconnected mapped drive isn't scanned twice. - media_roots = linux_run_media_mount_roots() - drive_roots = windows_drive_roots() + media_roots = [ + *external_media.linux_run_media_mount_roots(), + *external_media.macos_volume_roots(), + ] + drive_roots = external_media.windows_drive_roots() # Build once; the sandbox check and suggestion chips share it. allowed_roots = _build_browse_allowlist(media_roots, drive_roots) @@ -2034,19 +2059,19 @@ async def discard_remote_code_download( # Never delete a model that is loaded for inference. try: + from hub.services.models.deletion import _loaded_id_matches_repo from routes.inference import get_llama_cpp_backend + llama_backend = get_llama_cpp_backend() if llama_backend.is_loaded and llama_backend.model_identifier: - loaded = llama_backend.model_identifier.lower() - if loaded == model_name.lower() or loaded.startswith(model_name.lower()): + if _loaded_id_matches_repo(llama_backend.model_identifier, model_name): return {"deleted": False, "reason": "loaded"} except Exception: pass try: inference_backend = get_inference_backend() if inference_backend.active_model_name: - active = inference_backend.active_model_name.lower() - if active == model_name.lower() or active.startswith(model_name.lower()): + if _loaded_id_matches_repo(inference_backend.active_model_name, model_name): return {"deleted": False, "reason": "loaded"} except Exception: pass @@ -2588,13 +2613,10 @@ def _read_native_context_length(repo_id: str, is_local: bool) -> Optional[int]: if is_local: roots = [Path(repo_id)] else: - from huggingface_hub import constants as hf_constants - + from hub.utils.hf_cache_state import iter_repo_cache_dirs if not _is_valid_repo_id(repo_id): return None - cache_dir = Path(hf_constants.HF_HUB_CACHE) - target = f"models--{repo_id.replace('/', '--')}".lower() - roots = [e for e in cache_dir.iterdir() if e.name.lower() == target] + roots = list(iter_repo_cache_dirs("model", repo_id)) for root in roots: for f in _iter_gguf_paths(root): @@ -2623,18 +2645,15 @@ def _resolve_quant_gguf(repo_id: str, quant: str, is_local: bool) -> tuple[Optio if is_local: roots = [Path(repo_id)] else: - from huggingface_hub import constants as hf_constants + from hub.utils.hf_cache_state import iter_repo_cache_dirs if not _is_valid_repo_id(repo_id): return None, 0 - cache_dir = Path(hf_constants.HF_HUB_CACHE) - target = f"models--{repo_id.replace('/', '--')}".lower() roots = [] - for entry in cache_dir.iterdir(): - if entry.name.lower() == target: - snaps = entry / "snapshots" - if snaps.is_dir(): - roots.extend(s for s in snaps.iterdir() if s.is_dir()) + for entry in iter_repo_cache_dirs("model", repo_id): + snaps = entry / "snapshots" + if snaps.is_dir(): + roots.extend(s for s in snaps.iterdir() if s.is_dir()) want = _normalized_quant_label(quant) best_total = 0 @@ -2734,6 +2753,8 @@ async def get_gguf_variants( repo_id: str = Query( ..., description = "HuggingFace repo ID (e.g. 'unsloth/gemma-3-4b-it-GGUF')" ), + prefer_local_cache: bool = False, + local_path: Optional[str] = None, hf_token: Optional[str] = Query(None, description = "HuggingFace token for private repos"), hf_token_header: Optional[str] = Depends(get_hf_token), current_subject: str = Depends(get_current_subject), @@ -2745,9 +2766,16 @@ async def get_gguf_variants( response = await hub_gguf_variants.get_gguf_variants_response( repo_id, + prefer_local_cache = prefer_local_cache, + local_path = local_path, hf_token = hf_token, ) - local = is_local_path(repo_id) + context_model = ( + local_path + if prefer_local_cache and local_path and is_local_path(local_path) + else repo_id + ) + local = is_local_path(context_model) return GgufVariantsResponse( repo_id = response.repo_id, @@ -2769,7 +2797,7 @@ async def get_gguf_variants( # The header walk reads tokenizer arrays on dense models (tens of # ms per uncached file); keep it off the event loop. context_length = await asyncio.to_thread( - _read_native_context_length, repo_id, is_local = local + _read_native_context_length, context_model, is_local = local ), ) except HTTPException: @@ -2787,69 +2815,17 @@ async def get_gguf_download_progress( repo_id: str = Query(..., description = "HuggingFace repo ID"), variant: str = Query("", description = "Quantization variant (e.g. UD-TQ1_0)"), expected_bytes: int = Query(0, description = "Expected total download size in bytes"), + hf_token: Optional[str] = Depends(get_hf_token), current_subject: str = Depends(get_current_subject), ): - """Download progress from cached GGUF files for a specific variant. - - Tracks completed shards in snapshots and in-progress (.incomplete) - downloads in the blobs directory. - """ - try: - if not _is_valid_repo_id(repo_id): - return { - "downloaded_bytes": 0, - "expected_bytes": expected_bytes, - "progress": 0, - } - - from huggingface_hub import constants as hf_constants - - cache_dir = Path(hf_constants.HF_HUB_CACHE) - target = f"models--{repo_id.replace('/', '--')}".lower() - variant_lower = variant.lower().replace("-", "").replace("_", "") - downloaded_bytes = 0 - in_progress_bytes = 0 - for entry in cache_dir.iterdir(): - if entry.name.lower() == target: - # Completed .gguf files for this variant in snapshots. - # Exclude mmproj so a vision adapter can't satisfy a same-label - # main variant (e.g. mmproj-F16 vs an F16 weight). - for f in _iter_gguf_paths(entry): - if _is_mmproj_filename(f.name): - continue - rel = f.relative_to(entry).as_posix() - quant = _extract_quant_label(rel) - if _is_big_endian_gguf_path(rel, quant): - continue - rel_key = rel.lower().replace("-", "").replace("_", "") - if not variant_lower or variant_lower in rel_key: - try: - downloaded_bytes += f.stat().st_size - except OSError: - continue # broken symlink / unreadable: skip - # In-progress (.incomplete) downloads in blobs. - blobs_dir = entry / "blobs" - if blobs_dir.is_dir(): - for f in blobs_dir.iterdir(): - if f.is_file() and f.name.endswith(".incomplete"): - try: - in_progress_bytes += f.stat().st_size - except OSError: - continue - break - - total_progress_bytes = downloaded_bytes + in_progress_bytes - progress = min(total_progress_bytes / expected_bytes, 0.99) if expected_bytes > 0 else 0 - # Report 1.0 only when all bytes are in completed files. - if expected_bytes > 0 and downloaded_bytes >= expected_bytes: - progress = 1.0 - return { - "downloaded_bytes": total_progress_bytes, - "expected_bytes": expected_bytes, - "progress": round(progress, 3), - } - except Exception: - return {"downloaded_bytes": 0, "expected_bytes": expected_bytes, "progress": 0} + """Compatibility route backed by the shared multi-cache progress service.""" + from hub.services.models import downloads + return await downloads.get_gguf_download_progress_response( + repo_id, + variant = variant, + expected_bytes = expected_bytes, + hf_token = hf_token, + ) def _resolve_hf_cache_realpath(repo_dir: Path) -> Optional[str]: @@ -2874,98 +2850,12 @@ def _resolve_hf_cache_realpath(repo_dir: Path) -> Optional[str]: @router.get("/download-progress") async def get_download_progress( repo_id: str = Query(..., description = "HuggingFace repo ID"), + hf_token: Optional[str] = Depends(get_hf_token), current_subject: str = Depends(get_current_subject), ): - """Return download progress for any HuggingFace model repo. - - Checks the local HF cache for completed blobs and in-progress - (.incomplete) downloads. Gets the expected total size from the HF API - on the first call, then caches it for later polls. Also returns - ``cache_path``: the realpath of the snapshot dir (or cache repo root - if no snapshot yet) so the UI can show where weights live on disk. - """ - _empty = { - "downloaded_bytes": 0, - "expected_bytes": 0, - "progress": 0, - "cache_path": None, - } - try: - if not _is_valid_repo_id(repo_id): - return _empty - - from huggingface_hub import constants as hf_constants - - cache_dir = Path(hf_constants.HF_HUB_CACHE) - target = f"models--{repo_id.replace('/', '--')}".lower() - completed_bytes = 0 - in_progress_bytes = 0 - cache_path: Optional[str] = None - - for entry in cache_dir.iterdir(): - if entry.name.lower() != target: - continue - cache_path = _resolve_hf_cache_realpath(entry) - blobs_dir = entry / "blobs" - if not blobs_dir.is_dir(): - break - for f in blobs_dir.iterdir(): - if not f.is_file(): - continue - if f.name.endswith(".incomplete"): - in_progress_bytes += f.stat().st_size - else: - completed_bytes += f.stat().st_size - break - - downloaded_bytes = completed_bytes + in_progress_bytes - if downloaded_bytes == 0: - return {**_empty, "cache_path": cache_path} - - expected_bytes = _get_repo_size_cached(repo_id) - if expected_bytes <= 0: - # Total unknown; report bytes only, no percentage. - return { - "downloaded_bytes": downloaded_bytes, - "expected_bytes": 0, - "progress": 0, - "cache_path": cache_path, - } - - # 95% threshold (blob dedup can skew completed_bytes). Do NOT - # treat "no .incomplete files" as done: HF downloads sequentially, - # so none exist between files even when far from finished. - if completed_bytes >= expected_bytes * 0.95: - progress = 1.0 - else: - progress = min(downloaded_bytes / expected_bytes, 0.99) - return { - "downloaded_bytes": downloaded_bytes, - "expected_bytes": expected_bytes, - "progress": round(progress, 3), - "cache_path": cache_path, - } - except Exception as e: - logger.warning(f"Error checking download progress for {repo_id}: {e}") - return _empty - - -_repo_size_cache: dict[str, int] = {} - - -def _get_repo_size_cached(repo_id: str) -> int: - if repo_id in _repo_size_cache: - return _repo_size_cache[repo_id] - try: - from huggingface_hub import model_info as hf_model_info - - info = hf_model_info(repo_id, token = None, files_metadata = True) - total = sum(s.size for s in info.siblings if s.size) - _repo_size_cache[repo_id] = total - return total - except Exception as e: - logger.warning(f"Failed to get repo size for {repo_id}: {e}") - return 0 + """Compatibility route backed by the shared multi-cache progress service.""" + from hub.services.models import downloads + return await downloads.get_download_progress_response(repo_id, hf_token = hf_token) def _repo_in_any_hf_cache(model_name: str) -> bool: @@ -2978,25 +2868,13 @@ def _repo_in_any_hf_cache(model_name: str) -> bool: would delete a model they did not download via the scan. Mirrors the cache set in ``_all_hf_cache_scans`` but only probes for the one repo dir (cheap, no full scan). """ - from utils.paths import ( - hf_default_cache_dir, - legacy_hf_cache_dir, - resolve_cached_repo_id_case, - ) + from utils.paths import resolve_cached_repo_id_case dirname = f"models--{resolve_cached_repo_id_case(model_name).replace('/', '--')}" dirname_lower = dirname.lower() - candidates = [] - try: - from huggingface_hub.constants import HF_HUB_CACHE - candidates.append(Path(HF_HUB_CACHE)) - except Exception: - pass - for fn in (legacy_hf_cache_dir, hf_default_cache_dir): - try: - candidates.append(fn()) - except Exception: - continue + from hub.utils.hf_cache_state import hf_cache_roots + + candidates = hf_cache_roots() # resolve_cached_repo_id_case only normalizes the ACTIVE cache, but discard deletes # case-insensitively across all caches, so detect case-insensitively too -- else a # pre-existing case-variant repo is misreported as scan-created and deleted on decline. @@ -3020,38 +2898,8 @@ def _all_hf_cache_scans(): broken symlink, OS-redirected ~/.cache) is skipped, not fatal, so the Downloaded list never blanks out and downloads never leak into Recommended. """ - from huggingface_hub import scan_cache_dir - from utils.paths import legacy_hf_cache_dir, hf_default_cache_dir - - scans = [] - # Guard the active cache too: degrade to "no downloads" instead of raising. - try: - scans.append(scan_cache_dir()) - except Exception as exc: - logger.warning("Could not scan active HF cache: %s", exc) - - seen: set[str] = set() - try: - # Resolve the active cache dir for dedup. - from huggingface_hub.constants import HF_HUB_CACHE - seen.add(str(Path(HF_HUB_CACHE).resolve())) - except Exception: - pass - - for extra_fn in (legacy_hf_cache_dir, hf_default_cache_dir): - try: - extra = extra_fn() - # is_dir()/resolve() can raise on an inaccessible path; skip it. - if not extra.is_dir(): - continue - resolved = str(extra.resolve()) - if resolved in seen: - continue - seen.add(resolved) - scans.append(scan_cache_dir(cache_dir = str(extra))) - except Exception as exc: - logger.warning("Could not scan HF cache %s: %s", extra_fn.__name__, exc) - return scans + from hub.utils.inventory_scan import all_hf_cache_scans + return all_hf_cache_scans() def _is_gguf_filename(name: str) -> bool: @@ -3293,124 +3141,13 @@ async def list_cached_models( async def delete_cached_model( repo_id: str = Body(...), variant: Optional[str] = Body(None), + cache_path: Optional[str] = Body(None), + hf_token: Optional[str] = Depends(get_hf_token), current_subject: str = Depends(get_current_subject), ): - """Delete a cached model repo (or a specific GGUF variant) from the HF cache. - - With *variant*, only GGUF files matching that quant label are removed - (e.g. ``UD-Q4_K_XL``); otherwise the whole repo is deleted. Refuses - if the model is currently loaded for inference. - """ - if not _is_valid_repo_id(repo_id): - raise HTTPException(status_code = 400, detail = "Invalid repo_id format") - - # Refuse if the model is currently loaded. - try: - from routes.inference import get_llama_cpp_backend - llama_backend = get_llama_cpp_backend() - if llama_backend.is_loaded and llama_backend.model_identifier: - loaded_id = llama_backend.model_identifier.lower() - if loaded_id == repo_id.lower() or loaded_id.startswith(repo_id.lower()): - raise HTTPException( - status_code = 400, - detail = "Unload the model before deleting", - ) - except HTTPException: - raise - except Exception: - pass - - try: - inference_backend = get_inference_backend() - if inference_backend.active_model_name: - active = inference_backend.active_model_name.lower() - if active == repo_id.lower() or active.startswith(repo_id.lower()): - raise HTTPException( - status_code = 400, - detail = "Unload the model before deleting", - ) - except HTTPException: - raise - except Exception: - pass - - try: - cache_scans = _all_hf_cache_scans() - - target_repo = None - for hf_cache in cache_scans: - for repo_info in hf_cache.repos: - if repo_info.repo_type != "model": - continue - if repo_info.repo_id.lower() == repo_id.lower(): - target_repo = repo_info - break - if target_repo is not None: - break - - if target_repo is None: - raise HTTPException(status_code = 404, detail = "Model not found in cache") - - # โ”€โ”€ Per-variant GGUF deletion โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€ - if variant: - deleted_bytes = 0 - deleted_count = 0 - for rev in target_repo.revisions: - for f in rev.files: - if not _is_gguf_filename(f.file_name): - continue - quant = _extract_quant_label(f.file_name) - if quant.lower() != variant.lower(): - continue - # Delete the blob (data) and the snapshot symlink. - try: - blob = Path(f.blob_path) - snap = Path(f.file_path) - size = blob.stat().st_size if blob.exists() else 0 - if snap.exists() or snap.is_symlink(): - snap.unlink() - if blob.exists(): - blob.unlink() - deleted_bytes += size - deleted_count += 1 - except Exception as e: - logger.warning(f"Failed to delete {f.file_name}: {e}") - - if deleted_count == 0: - raise HTTPException( - status_code = 404, - detail = f"Variant {variant} not found in cache for {repo_id}", - ) - - freed_mb = deleted_bytes / (1024 * 1024) - logger.info( - f"Deleted {deleted_count} file(s) for {repo_id} variant {variant}: " - f"{freed_mb:.1f} MB freed" - ) - return {"status": "deleted", "repo_id": repo_id, "variant": variant} - - # โ”€โ”€ Full repo deletion โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€ - revision_hashes = [rev.commit_hash for rev in target_repo.revisions] - if not revision_hashes: - raise HTTPException(status_code = 404, detail = "No revisions found for model") - - delete_strategy = hf_cache.delete_revisions(*revision_hashes) - logger.info( - f"Deleting cached model {repo_id}: " - f"{delete_strategy.expected_freed_size_str} will be freed" - ) - delete_strategy.execute() - - return {"status": "deleted", "repo_id": repo_id} - - except HTTPException: - raise - except Exception as e: - logger.error(f"Error deleting cached model {repo_id}: {e}", exc_info = True) - raise HTTPException( - status_code = 500, - detail = "Failed to delete cached model", - ) + """Compatibility route backed by the shared multi-cache deletion service.""" + from hub.services.models import deletion + return await deletion.delete_cached_model_response(repo_id, variant, hf_token, cache_path) def _resolve_cached_model_path(repo_id: str, variant: Optional[str]) -> Path: diff --git a/studio/backend/routes/settings.py b/studio/backend/routes/settings.py index f36c8870e3..fef18a9145 100644 --- a/studio/backend/routes/settings.py +++ b/studio/backend/routes/settings.py @@ -60,6 +60,7 @@ from utils.embedding_model_settings import ( set_rag_embedding_model, validate_embedding_model, ) +from utils.hf_cache_settings import cache_status, get_hf_cache_paths, set_hf_cache_home router = APIRouter() @@ -89,6 +90,23 @@ class HelperPrecacheResponse(BaseModel): disabled_by_env: bool +class HuggingFaceCachePayload(BaseModel): + cache_home: Optional[str] = Field(default = None, max_length = 4096) + + +class HuggingFaceCacheResponse(BaseModel): + cache_home: str + hub_cache: str + xet_cache: str + source: Literal["default", "studio", "environment"] + editable: bool + is_custom: bool + available: bool + writable: bool + free_bytes: Optional[int] = None + environment_variable: Optional[str] = None + + class OpenAIAutoSwitchPayload(BaseModel): enabled: bool # None leaves the stored value untouched (partial updates can't clobber it). @@ -135,6 +153,30 @@ def _helper_precache_response(enabled: bool | None = None) -> HelperPrecacheResp ) +def _hugging_face_cache_response() -> HuggingFaceCacheResponse: + return HuggingFaceCacheResponse(**cache_status(get_hf_cache_paths())) + + +@router.get("/hugging-face-cache", response_model = HuggingFaceCacheResponse) +def get_hugging_face_cache( + current_subject: str = Depends(get_current_subject), +) -> HuggingFaceCacheResponse: + return _hugging_face_cache_response() + + +@router.put("/hugging-face-cache", response_model = HuggingFaceCacheResponse) +def update_hugging_face_cache( + payload: HuggingFaceCachePayload, current_subject: str = Depends(get_current_subject) +) -> HuggingFaceCacheResponse: + try: + set_hf_cache_home(payload.cache_home) + except RuntimeError as exc: + raise HTTPException(status_code = 409, detail = str(exc)) from exc + except ValueError as exc: + raise HTTPException(status_code = 400, detail = str(exc)) from exc + return _hugging_face_cache_response() + + @router.get("/upload-limit", response_model = UploadLimitResponse) def get_upload_limit(current_subject: str = Depends(get_current_subject)) -> UploadLimitResponse: return _upload_limit_response(get_upload_limit_mb()) diff --git a/studio/backend/tests/test_cached_gguf_routes.py b/studio/backend/tests/test_cached_gguf_routes.py index b3e6255d55..8ee72259f4 100644 --- a/studio/backend/tests/test_cached_gguf_routes.py +++ b/studio/backend/tests/test_cached_gguf_routes.py @@ -66,6 +66,66 @@ def test_iter_gguf_paths_matches_extension_case_insensitively(tmp_path): assert result == ["Q4_K_M.gguf", "Q8_0.GGUF"] +def test_legacy_hf_scan_uses_snapshot_path_for_inactive_cache(tmp_path): + repo = tmp_path / "models--Org--Model" + snapshot = repo / "snapshots" / "revision" + snapshot.mkdir(parents = True) + + [row] = models_route._scan_hf_cache(tmp_path, active_cache = False) + + assert row.model_id == "Org/Model" + assert row.id == str(snapshot.resolve()) + assert row.path == str(snapshot.resolve()) + + +def test_collect_local_models_scans_previous_cache(monkeypatch, tmp_path): + active = tmp_path / "active" + previous = tmp_path / "previous" + active.mkdir() + snapshot = previous / "models--Org--Previous" / "snapshots" / "revision" + snapshot.mkdir(parents = True) + + monkeypatch.setattr(models_route, "_resolve_hf_cache_dir", lambda: active) + monkeypatch.setattr("utils.paths.legacy_hf_cache_dir", lambda: tmp_path / "legacy") + monkeypatch.setattr("utils.paths.hf_default_cache_dir", lambda: tmp_path / "default") + monkeypatch.setattr("utils.paths.lmstudio_model_dirs", lambda: []) + monkeypatch.setattr("utils.hf_cache_settings.known_hf_hub_caches", lambda: [active, previous]) + monkeypatch.setattr("storage.studio_db.list_scan_folders", lambda: []) + + rows = models_route.collect_local_models(tmp_path / "models") + + previous_row = next(row for row in rows if row.model_id == "Org/Previous") + assert previous_row.id == str(snapshot.resolve()) + + +def test_collect_local_models_prefers_complete_previous_copy(monkeypatch, tmp_path): + active = tmp_path / "active" + previous = tmp_path / "previous" + active_partial = active / "models--Org--Model" / "blobs" / "abc.incomplete" + active_partial.parent.mkdir(parents = True) + active_partial.write_bytes(b"partial") + snapshot = previous / "models--Org--Model" / "snapshots" / "revision" + snapshot.mkdir(parents = True) + (snapshot / "model.safetensors").write_bytes(b"complete") + + monkeypatch.setattr(models_route, "_resolve_hf_cache_dir", lambda: active) + monkeypatch.setattr("utils.paths.legacy_hf_cache_dir", lambda: tmp_path / "legacy") + monkeypatch.setattr("utils.paths.hf_default_cache_dir", lambda: tmp_path / "default") + monkeypatch.setattr("utils.paths.lmstudio_model_dirs", lambda: []) + monkeypatch.setattr( + "utils.hf_cache_settings.known_hf_hub_caches", + lambda: [active, previous], + ) + monkeypatch.setattr("storage.studio_db.list_scan_folders", lambda: []) + + rows = models_route.collect_local_models(tmp_path / "models") + + [row] = [row for row in rows if row.model_id == "Org/Model"] + assert row.id == str(snapshot.resolve()) + assert row.partial is False + assert row.active_cache is False + + def test_list_cached_gguf_includes_non_suffix_repo_when_cache_contains_gguf(monkeypatch, tmp_path): repo = _repo( "HauhauCS/Gemma-4-E4B-Uncensored-HauhauCS-Aggressive", @@ -573,33 +633,14 @@ def _gfile(name: str, size: int, mtime: float) -> SimpleNamespace: ) -def test_all_hf_cache_scans_survives_inaccessible_aux_cache(monkeypatch, tmp_path): - """An unreadable auxiliary cache (e.g. an inaccessible - ``~/.cache/huggingface/hub``) must be skipped, not abort the scan. - Regression guard for ``extra.is_dir()`` raising and wiping the response. - """ - import huggingface_hub - import utils.paths as paths_mod +def test_all_hf_cache_scans_uses_shared_inventory(monkeypatch, tmp_path): + from hub.utils import inventory_scan active = SimpleNamespace( repos = [_repo("Org/Active", [_file("Q4_K_M.gguf", 5_000)], tmp_path / "active")] ) - def _fake_scan(cache_dir = None): - if cache_dir is None: - return active - raise AssertionError("auxiliary scan should have been skipped") - - class _Boom: - def is_dir(self): - raise PermissionError(13, "Permission denied") - - def resolve(self): - raise PermissionError(13, "Permission denied") - - monkeypatch.setattr(huggingface_hub, "scan_cache_dir", _fake_scan) - monkeypatch.setattr(paths_mod, "legacy_hf_cache_dir", lambda: _Boom()) - monkeypatch.setattr(paths_mod, "hf_default_cache_dir", lambda: _Boom()) + monkeypatch.setattr(inventory_scan, "all_hf_cache_scans", lambda: [active]) scans = models_route._all_hf_cache_scans() assert scans == [active] @@ -686,13 +727,17 @@ def test_gguf_variants_mmproj_does_not_mark_quant_downloaded(monkeypatch, tmp_pa "list_gguf_variants", lambda repo_id, hf_token = None: (variants, True, []), ) - monkeypatch.setattr(GV, "_local_main_gguf_blobs_by_quant", lambda _repo_id: {}) + monkeypatch.setattr( + GV, + "_local_main_gguf_blobs_by_quant", + lambda _repo_id, repo_cache_dir = None: {}, + ) snap = tmp_path / "models--org--repo" / "snapshots" / "rev" snap.mkdir(parents = True) (snap / "model-Q4_K_M.gguf").write_bytes(b"x" * 10_000) # real weight, fully present (snap / "mmproj-F16.gguf").write_bytes(b"y" * 20_000) # mmproj adapter, label "F16" - monkeypatch.setattr(GV, "iter_hf_cache_snapshots", lambda _repo_id: [snap]) + monkeypatch.setattr(GV, "iter_hf_cache_snapshots", lambda _repo_id, root = None: [snap]) result = asyncio.run( models_route.get_gguf_variants( @@ -705,6 +750,52 @@ def test_gguf_variants_mmproj_does_not_mark_quant_downloaded(monkeypatch, tmp_pa assert flags["F16"] is False +def test_gguf_variants_route_scopes_local_probe_to_selected_cache(monkeypatch, tmp_path): + snapshot = tmp_path / "inactive" / "models--org--repo" / "snapshots" / "rev" + snapshot.mkdir(parents = True) + calls = [] + + async def scoped_variants(repo_id, **kwargs): + calls.append((repo_id, kwargs)) + return SimpleNamespace( + repo_id = repo_id, + variants = [], + has_vision = False, + default_variant = None, + ) + + context_calls = [] + monkeypatch.setattr(GV, "get_gguf_variants_response", scoped_variants) + monkeypatch.setattr( + models_route, + "_read_native_context_length", + lambda model, *, is_local: context_calls.append((model, is_local)) or 8192, + ) + + result = asyncio.run( + models_route.get_gguf_variants( + repo_id = "org/repo", + prefer_local_cache = True, + local_path = str(snapshot), + hf_token = None, + current_subject = "test-user", + ) + ) + + assert calls == [ + ( + "org/repo", + { + "prefer_local_cache": True, + "local_path": str(snapshot), + "hf_token": None, + }, + ) + ] + assert context_calls == [(str(snapshot), True)] + assert result.context_length == 8192 + + def test_gguf_variants_ignore_big_endian_siblings(monkeypatch, tmp_path): siblings = [ SimpleNamespace(rfilename = "model-Q4_K_M-be.gguf", size = 100), @@ -726,12 +817,16 @@ def test_gguf_variants_ignore_big_endian_siblings(monkeypatch, tmp_path): siblings, ), ) - monkeypatch.setattr(GV, "_local_main_gguf_blobs_by_quant", lambda _repo_id: {}) + monkeypatch.setattr( + GV, + "_local_main_gguf_blobs_by_quant", + lambda _repo_id, repo_cache_dir = None: {}, + ) snap = tmp_path / "models--org--repo" / "snapshots" / "rev" snap.mkdir(parents = True) (snap / "model-Q4_K_M.gguf").write_bytes(b"x" * 10) - monkeypatch.setattr(GV, "iter_hf_cache_snapshots", lambda _repo_id: [snap]) + monkeypatch.setattr(GV, "iter_hf_cache_snapshots", lambda _repo_id, root = None: [snap]) result = asyncio.run( models_route.get_gguf_variants( @@ -758,12 +853,16 @@ def test_gguf_variants_cached_big_endian_does_not_satisfy_variant(monkeypatch, t "list_gguf_variants", lambda repo_id, hf_token = None: (variants, False, []), ) - monkeypatch.setattr(GV, "_local_main_gguf_blobs_by_quant", lambda _repo_id: {}) + monkeypatch.setattr( + GV, + "_local_main_gguf_blobs_by_quant", + lambda _repo_id, repo_cache_dir = None: {}, + ) snap = tmp_path / "models--org--repo" / "snapshots" / "rev" snap.mkdir(parents = True) (snap / "model-Q4_K_M-be.gguf").write_bytes(b"x" * 10) - monkeypatch.setattr(GV, "iter_hf_cache_snapshots", lambda _repo_id: [snap]) + monkeypatch.setattr(GV, "iter_hf_cache_snapshots", lambda _repo_id, root = None: [snap]) result = asyncio.run( models_route.get_gguf_variants( @@ -774,66 +873,82 @@ def test_gguf_variants_cached_big_endian_does_not_satisfy_variant(monkeypatch, t assert result.variants[0].downloaded is False -def test_gguf_download_progress_excludes_mmproj(monkeypatch, tmp_path): - """A cached mmproj adapter must not count toward a same-label main - variant's download progress (mmproj-F16 vs an F16 weight).""" - import huggingface_hub.constants as hf_constants +def test_legacy_gguf_progress_delegates_to_shared_service(monkeypatch): + calls = [] - monkeypatch.setattr(hf_constants, "HF_HUB_CACHE", str(tmp_path)) - snap = tmp_path / "models--org--repo" / "snapshots" / "rev" - snap.mkdir(parents = True) - (snap / "mmproj-F16.gguf").write_bytes(b"y" * 20_000) # only the adapter on disk + async def shared(repo_id, *, variant, expected_bytes, hf_token): + calls.append((repo_id, variant, expected_bytes, hf_token)) + return {"downloaded_bytes": 10, "expected_bytes": 20, "progress": 0.5} - result = asyncio.run( - models_route.get_gguf_download_progress( - repo_id = "org/repo", - variant = "F16", - expected_bytes = 20_000, - current_subject = "test-user", - ) + monkeypatch.setattr( + "hub.services.models.downloads.get_gguf_download_progress_response", + shared, ) - assert result["downloaded_bytes"] == 0 - assert result["progress"] == 0 - - -def test_gguf_download_progress_excludes_big_endian_sibling(monkeypatch, tmp_path): - import huggingface_hub.constants as hf_constants - - monkeypatch.setattr(hf_constants, "HF_HUB_CACHE", str(tmp_path)) - snap = tmp_path / "models--org--repo" / "snapshots" / "rev" - snap.mkdir(parents = True) - (snap / "model-Q4_K_M-be.gguf").write_bytes(b"y" * 20_000) - result = asyncio.run( models_route.get_gguf_download_progress( repo_id = "org/repo", variant = "Q4_K_M", - expected_bytes = 20_000, + expected_bytes = 20, + hf_token = "token", current_subject = "test-user", ) ) - assert result["downloaded_bytes"] == 0 - assert result["progress"] == 0 + assert result["progress"] == 0.5 + assert calls == [("org/repo", "Q4_K_M", 20, "token")] -def test_gguf_download_progress_counts_quant_subdir(monkeypatch, tmp_path): - import huggingface_hub.constants as hf_constants +def test_legacy_model_progress_delegates_to_shared_service(monkeypatch): + calls = [] - monkeypatch.setattr(hf_constants, "HF_HUB_CACHE", str(tmp_path)) - snap = tmp_path / "models--org--repo" / "snapshots" / "rev" / "Q4_K_M" - snap.mkdir(parents = True) - (snap / "foo.gguf").write_bytes(b"x" * 20_000) + async def shared(repo_id, *, hf_token): + calls.append((repo_id, hf_token)) + return {"downloaded_bytes": 10, "expected_bytes": 20, "progress": 0.5} + + monkeypatch.setattr( + "hub.services.models.downloads.get_download_progress_response", + shared, + ) result = asyncio.run( - models_route.get_gguf_download_progress( + models_route.get_download_progress( repo_id = "org/repo", - variant = "Q4_K_M", - expected_bytes = 20_000, + hf_token = "token", current_subject = "test-user", ) ) - assert result["downloaded_bytes"] == 20_000 - assert result["progress"] == 1.0 + assert result["progress"] == 0.5 + assert calls == [("org/repo", "token")] + + +def test_legacy_delete_delegates_to_shared_service(monkeypatch): + calls = [] + + async def shared( + repo_id, + variant, + hf_token, + cache_path = None, + ): + calls.append((repo_id, variant, hf_token, cache_path)) + return {"status": "deleted", "repo_id": repo_id} + + monkeypatch.setattr( + "hub.services.models.deletion.delete_cached_model_response", + shared, + ) + + result = asyncio.run( + models_route.delete_cached_model( + repo_id = "org/repo", + variant = None, + cache_path = "/data/hf/hub", + hf_token = "token", + current_subject = "test-user", + ) + ) + + assert result == {"status": "deleted", "repo_id": "org/repo"} + assert calls == [("org/repo", None, "token", "/data/hf/hub")] diff --git a/studio/backend/tests/test_consent_gate.py b/studio/backend/tests/test_consent_gate.py index 804221ec7e..181e0c9fad 100644 --- a/studio/backend/tests/test_consent_gate.py +++ b/studio/backend/tests/test_consent_gate.py @@ -873,6 +873,7 @@ class TestScannerCoversAllExecutableCode: repo, fn, token = None, + cache_dir = None, ): if fn == "config.json": import json @@ -899,6 +900,7 @@ class TestScannerCoversAllExecutableCode: repo, fn, token = None, + cache_dir = None, ): import json import tempfile @@ -932,6 +934,7 @@ class TestScannerCoversAllExecutableCode: repo, fn, token = None, + cache_dir = None, ): import json import tempfile @@ -972,6 +975,7 @@ class TestScannerCoversAllExecutableCode: repo, fn, token = None, + cache_dir = None, ): import json import tempfile @@ -1008,6 +1012,7 @@ class TestScannerCoversAllExecutableCode: repo, fn, token = None, + cache_dir = None, ): import json import tempfile @@ -1037,6 +1042,7 @@ class TestScannerCoversAllExecutableCode: repo, fn, token = None, + cache_dir = None, ): import json import tempfile @@ -1079,6 +1085,7 @@ class TestScannerCoversAllExecutableCode: repo, fn, token = None, + cache_dir = None, ): import json import tempfile @@ -1120,6 +1127,7 @@ class TestScannerCoversAllExecutableCode: repo, fn, token = None, + cache_dir = None, ): import json import tempfile @@ -1182,6 +1190,7 @@ class TestScannerCoversAllExecutableCode: repo, fn, token = None, + cache_dir = None, ): import json import tempfile diff --git a/studio/backend/tests/test_gguf_load_cache_reuse.py b/studio/backend/tests/test_gguf_load_cache_reuse.py index 62596fcc8a..6d1fac980b 100644 --- a/studio/backend/tests/test_gguf_load_cache_reuse.py +++ b/studio/backend/tests/test_gguf_load_cache_reuse.py @@ -103,6 +103,10 @@ def _build_cache( @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 @@ -117,6 +121,61 @@ def _fail_get_paths_info(*_args, **_kwargs): 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() diff --git a/studio/backend/tests/test_hf_cache_settings.py b/studio/backend/tests/test_hf_cache_settings.py new file mode 100644 index 0000000000..1875d61809 --- /dev/null +++ b/studio/backend/tests/test_hf_cache_settings.py @@ -0,0 +1,290 @@ +# SPDX-License-Identifier: AGPL-3.0-only +# Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. See /studio/LICENSE.AGPL-3.0 + +from __future__ import annotations + +import os +import sys +import threading +import time +from pathlib import Path + +import pytest + +_BACKEND_DIR = str(Path(__file__).resolve().parent.parent) +if _BACKEND_DIR not in sys.path: + sys.path.insert(0, _BACKEND_DIR) + +from hub.services.models.common import _local_model_info +from utils import hf_cache_settings +from utils import native_path_leases + + +@pytest.fixture() +def settings_store(monkeypatch, tmp_path): + store = {} + monkeypatch.setattr(hf_cache_settings, "_EXPLICIT_CACHE_ENV", {}) + monkeypatch.setenv("XDG_CACHE_HOME", str(tmp_path / "xdg")) + monkeypatch.setattr( + "storage.studio_db.get_app_setting", + lambda key, fallback = None: store.get(key, fallback), + ) + monkeypatch.setattr( + "storage.studio_db.upsert_app_settings", + lambda values: store.update(values) or values, + ) + return store + + +def test_studio_cache_switch_is_live_and_keeps_history(settings_store, tmp_path): + first = tmp_path / "external-a" / "huggingface" + second = tmp_path / "external-b" / "huggingface" + first.parent.mkdir() + second.parent.mkdir() + + selected = hf_cache_settings.set_hf_cache_home(str(first)) + assert selected.hub_cache == first / "hub" + assert selected.xet_cache == first / "xet" + assert selected.child_env({}) == { + "HF_HUB_CACHE": str(first / "hub"), + "HF_XET_CACHE": str(first / "xet"), + } + + hf_cache_settings.set_hf_cache_home(str(second)) + assert settings_store[hf_cache_settings.CACHE_HISTORY_SETTING_KEY] == [str(first)] + assert first / "hub" in hf_cache_settings.known_hf_hub_caches() + + reset = hf_cache_settings.set_hf_cache_home(None) + assert reset.source == "default" + assert second in hf_cache_settings.known_hf_cache_homes() + + +def test_environment_cache_is_read_only(monkeypatch, tmp_path): + custom = tmp_path / "managed" + monkeypatch.setattr( + hf_cache_settings, + "_EXPLICIT_CACHE_ENV", + {"HF_HOME": str(custom)}, + ) + paths = hf_cache_settings.get_hf_cache_paths() + assert paths.source == "environment" + assert paths.editable is False + assert paths.hub_cache == custom / "hub" + with pytest.raises(RuntimeError, match = "environment variable"): + hf_cache_settings.set_hf_cache_home(str(tmp_path / "other")) + + +def test_explicit_hub_cache_is_the_displayed_location(monkeypatch, tmp_path): + custom_hub = tmp_path / "models-cache" + custom_hub.mkdir() + monkeypatch.setattr( + hf_cache_settings, + "_EXPLICIT_CACHE_ENV", + {"HF_HUB_CACHE": str(custom_hub)}, + ) + + paths = hf_cache_settings.get_hf_cache_paths() + status = hf_cache_settings.cache_status(paths) + + assert paths.cache_home == custom_hub + assert paths.hub_cache == custom_hub + assert status["cache_home"] == str(custom_hub) + assert status["available"] is True + assert custom_hub / "hub" not in hf_cache_settings.known_hf_hub_caches() + + +def test_explicit_hub_cache_display_wins_over_hf_home(monkeypatch, tmp_path): + hf_home = tmp_path / "hf-home" + custom_hub = tmp_path / "other-disk" / "models-cache" + hf_home.mkdir() + custom_hub.mkdir(parents = True) + monkeypatch.setattr( + hf_cache_settings, + "_EXPLICIT_CACHE_ENV", + {"HF_HOME": str(hf_home), "HF_HUB_CACHE": str(custom_hub)}, + ) + + paths = hf_cache_settings.get_hf_cache_paths() + + assert paths.cache_home == custom_hub + assert paths.hub_cache == custom_hub + assert paths.xet_cache == hf_home / "xet" + assert custom_hub / "hub" not in hf_cache_settings.known_hf_hub_caches() + assert hf_home / "hub" in hf_cache_settings.known_hf_hub_caches() + + +def test_xet_only_override_keeps_model_cache_editable(settings_store, monkeypatch, tmp_path): + xet_cache = tmp_path / "chunks" + stored = tmp_path / "stored-cache" + settings_store[hf_cache_settings.CACHE_HOME_SETTING_KEY] = str(stored) + monkeypatch.setattr( + hf_cache_settings, + "_EXPLICIT_CACHE_ENV", + {"HF_XET_CACHE": str(xet_cache)}, + ) + + paths = hf_cache_settings.get_hf_cache_paths() + + assert paths.cache_home == stored + assert paths.hub_cache == stored / "hub" + assert paths.xet_cache == xet_cache + assert paths.editable is True + + selected = tmp_path / "selected-cache" + selected.parent.mkdir(exist_ok = True) + updated = hf_cache_settings.set_hf_cache_home(str(selected)) + assert updated.hub_cache == selected / "hub" + assert updated.xet_cache == xet_cache + + +def test_worker_environment_is_applied_before_import(monkeypatch, tmp_path): + hub = str(tmp_path / "hub") + xet = str(tmp_path / "xet") + observed = {} + + class Module: + @staticmethod + def run(): + import os + return os.environ["HF_HUB_CACHE"], os.environ["HF_XET_CACHE"] + + def fake_import(name): + import os + + observed["name"] = name + observed["hub"] = os.environ.get("HF_HUB_CACHE") + return Module + + monkeypatch.setattr(native_path_leases.importlib, "import_module", fake_import) + result = native_path_leases.run_without_native_path_secret( + "fake.worker", + "run", + {"HF_HUB_CACHE": hub, "HF_XET_CACHE": xet}, + ) + assert observed == {"name": "fake.worker", "hub": hub} + assert result == (hub, xet) + + +def test_spawn_environment_is_applied_then_restored(monkeypatch, tmp_path): + hub = str(tmp_path / "hub") + xet = str(tmp_path / "xet") + monkeypatch.setenv("HF_HUB_CACHE", "parent-hub") + monkeypatch.delenv("HF_XET_CACHE", raising = False) + + with hf_cache_settings.child_environment_for_spawn({"HF_HUB_CACHE": hub, "HF_XET_CACHE": xet}): + import os + assert os.environ["HF_HUB_CACHE"] == hub + assert os.environ["HF_XET_CACHE"] == xet + + assert os.environ["HF_HUB_CACHE"] == "parent-hub" + assert "HF_XET_CACHE" not in os.environ + + +def test_spawn_environment_supports_nested_contexts(monkeypatch): + monkeypatch.setenv("HF_HUB_CACHE", "parent") + + with hf_cache_settings.child_environment_for_spawn({"HF_HUB_CACHE": "outer"}): + assert os.environ["HF_HUB_CACHE"] == "outer" + with hf_cache_settings.child_environment_for_spawn({"HF_HUB_CACHE": "inner"}): + assert os.environ["HF_HUB_CACHE"] == "inner" + assert os.environ["HF_HUB_CACHE"] == "outer" + + assert os.environ["HF_HUB_CACHE"] == "parent" + + +def test_spawn_environment_serializes_threads(monkeypatch): + monkeypatch.setenv("HF_HUB_CACHE", "parent") + first_entered = threading.Event() + release_first = threading.Event() + observations: list[tuple[str, str]] = [] + + def first(): + with hf_cache_settings.child_environment_for_spawn({"HF_HUB_CACHE": "first"}): + observations.append(("first", os.environ["HF_HUB_CACHE"])) + first_entered.set() + assert release_first.wait(timeout = 2) + + def second(): + assert first_entered.wait(timeout = 2) + with hf_cache_settings.child_environment_for_spawn({"HF_HUB_CACHE": "second"}): + observations.append(("second", os.environ["HF_HUB_CACHE"])) + + first_thread = threading.Thread(target = first) + second_thread = threading.Thread(target = second) + first_thread.start() + second_thread.start() + assert first_entered.wait(timeout = 2) + time.sleep(0.02) + assert observations == [("first", "first")] + release_first.set() + first_thread.join(timeout = 2) + second_thread.join(timeout = 2) + + assert observations == [("first", "first"), ("second", "second")] + assert os.environ["HF_HUB_CACHE"] == "parent" + + +def test_cache_switch_invalidates_inventory(settings_store, tmp_path, monkeypatch): + invalidations = [] + monkeypatch.setattr( + "hub.utils.inventory_scan.invalidate_hf_cache_scans", + lambda: invalidations.append(True), + ) + selected = tmp_path / "external" / "huggingface" + selected.parent.mkdir() + + hf_cache_settings.set_hf_cache_home(str(selected)) + + assert invalidations == [True] + + +def test_cache_validation_write_tests_hub_and_xet(settings_store, tmp_path, monkeypatch): + selected = tmp_path / "external" / "huggingface" + selected.parent.mkdir() + tested = [] + real_named_temporary_file = hf_cache_settings.tempfile.NamedTemporaryFile + + def recording_write_test(*args, **kwargs): + tested.append(Path(kwargs["dir"])) + return real_named_temporary_file(*args, **kwargs) + + monkeypatch.setattr( + hf_cache_settings.tempfile, + "NamedTemporaryFile", + recording_write_test, + ) + + hf_cache_settings.set_hf_cache_home(str(selected)) + + assert tested == [selected / "hub", selected / "xet"] + + +def test_cache_validation_rejects_unwritable_child(settings_store, tmp_path, monkeypatch): + selected = tmp_path / "external" / "huggingface" + selected.parent.mkdir() + + def reject_hub(*args, **kwargs): + if Path(kwargs["dir"]).name == "hub": + raise PermissionError("read-only") + raise AssertionError("xet should not be tested after hub fails") + + monkeypatch.setattr(hf_cache_settings.tempfile, "NamedTemporaryFile", reject_hub) + + with pytest.raises(ValueError, match = "permission"): + hf_cache_settings.set_hf_cache_home(str(selected)) + + +def test_inactive_cache_model_loads_from_snapshot_path(tmp_path): + snapshot = tmp_path / "snapshots" / "revision" + snapshot.mkdir(parents = True) + row = _local_model_info( + scan_path = snapshot, + load_path = snapshot, + source = "hf_cache", + model_format = "safetensors", + model_id = "org/model", + active_cache = False, + ) + assert row.model_id == "org/model" + assert row.active_cache is False + assert row.load_id == str(snapshot) diff --git a/studio/backend/tests/test_hf_xet_fallback.py b/studio/backend/tests/test_hf_xet_fallback.py index 48aff29659..a037ea2579 100644 --- a/studio/backend/tests/test_hf_xet_fallback.py +++ b/studio/backend/tests/test_hf_xet_fallback.py @@ -101,13 +101,23 @@ def test_shim_injects_studio_prepare_on_http_retry(monkeypatch): prepared = [] monkeypatch.setattr( "hub.utils.download_registry.prepare_cache_for_transport", - lambda repo_type, repo_id, mode, *a, **k: prepared.append((repo_type, repo_id, mode)), + lambda repo_type, repo_id, mode, *a, **k: prepared.append( + (repo_type, repo_id, mode, k.get("root")) + ), ) - out = xf.hf_hub_download_with_xet_fallback(DL_REPO, FILE, None) + selected_cache = "/captured/hub" + out = xf.hf_hub_download_with_xet_fallback( + DL_REPO, + FILE, + None, + cache_dir = selected_cache, + ) assert out == "/cache/model.gguf" assert seen_disable_xet == [False, True] # Xet first, then HTTP - assert prepared == [("model", DL_REPO, "http")], "shim must run Unsloth's marker-aware prep" + assert prepared == [ + ("model", DL_REPO, "http", Path(selected_cache)) + ], "shim must prepare the cache captured by the download" def test_shim_snapshot_injects_studio_prepare(monkeypatch): @@ -120,10 +130,22 @@ def test_shim_snapshot_injects_studio_prepare(monkeypatch): return "/tmp/snap-dir" monkeypatch.setattr(xf, "_shared_snapshot_download_with_xet_fallback", fake_snapshot) - out = xf.snapshot_download_with_xet_fallback("org/model") + selected_cache = "/captured/hub" + out = xf.snapshot_download_with_xet_fallback( + "org/model", + cache_dir = selected_cache, + ) assert out == "/tmp/snap-dir" assert captured["repo_id"] == "org/model" - assert captured["prepare_for_http_fn"] is xf._studio_prepare_for_http + prepared = [] + monkeypatch.setattr( + "hub.utils.download_registry.prepare_cache_for_transport", + lambda repo_type, repo_id, mode, *a, **k: prepared.append( + (repo_type, repo_id, mode, k.get("root")) + ), + ) + captured["prepare_for_http_fn"]("model", "org/model") + assert prepared == [("model", "org/model", "http", Path(selected_cache))] def test_degrades_gracefully_without_shared_helper(monkeypatch): diff --git a/studio/backend/tests/test_linux_external_media_paths.py b/studio/backend/tests/test_linux_external_media_paths.py index b735bd1132..8373cdd6bb 100644 --- a/studio/backend/tests/test_linux_external_media_paths.py +++ b/studio/backend/tests/test_linux_external_media_paths.py @@ -254,8 +254,10 @@ def test_legacy_browse_allowlist_includes_linux_run_media_mounts(monkeypatch, tm ) fake_external_media = SimpleNamespace( linux_run_media_mount_roots = lambda: [media_root], + macos_volume_roots = lambda: [], windows_drive_roots = lambda: [], ) + fake_paths.external_media = fake_external_media fake_studio_db = SimpleNamespace( list_scan_folders = lambda: [], contains_sensitive_path_component = studio_db.contains_sensitive_path_component, diff --git a/studio/backend/tests/test_model_update_robustness.py b/studio/backend/tests/test_model_update_robustness.py index 7d98766616..7e5d793740 100644 --- a/studio/backend/tests/test_model_update_robustness.py +++ b/studio/backend/tests/test_model_update_robustness.py @@ -112,13 +112,21 @@ def patch_hub_gguf(monkeypatch): blob_ids = [local_blob], gguf_files = {"model-Q4_K_M.gguf": 1000}, ) + monkeypatch.setattr( + "utils.hf_cache_settings.get_hf_cache_paths", + lambda: SimpleNamespace(hub_cache = tmp_path), + ) monkeypatch.setattr( GV, "list_gguf_variants", lambda r, hf_token = None: (_variants(), False, [remote_sibling]), raising = True, ) - monkeypatch.setattr(GV, "iter_hf_cache_snapshots", lambda _repo_id: [snap]) + monkeypatch.setattr( + GV, + "iter_hf_cache_snapshots", + lambda _repo_id, root = None: [snap], + ) monkeypatch.setattr( CI, "all_hf_cache_scans", @@ -217,6 +225,10 @@ def test_variant_update_check_detects_companion_only_update( companion_path: 100, }, ) + monkeypatch.setattr( + "utils.hf_cache_settings.get_hf_cache_paths", + lambda: SimpleNamespace(hub_cache = tmp_path), + ) siblings = [ patch_hub_gguf.sibling("model-Q4_K_M.gguf", 1000, "mainsha"), patch_hub_gguf.sibling(companion_path, 100, "new-companion"), @@ -227,7 +239,11 @@ def test_variant_update_check_detects_companion_only_update( lambda r, hf_token = None: (_variants(), has_vision, siblings), raising = True, ) - monkeypatch.setattr(GV, "iter_hf_cache_snapshots", lambda _repo_id: [snap]) + monkeypatch.setattr( + GV, + "iter_hf_cache_snapshots", + lambda _repo_id, root = None: [snap], + ) monkeypatch.setattr( CI, "all_hf_cache_scans", @@ -372,7 +388,7 @@ def test_cached_gguf_scan_keeps_download_timestamp(monkeypatch, tmp_path): monkeypatch.setattr( CI, "_gguf_variant_state_summary", - lambda _repo_id: (False, 0), + lambda _repo_id, **_kwargs: (False, 0), ) rows = CI._scan_cached_gguf() @@ -630,7 +646,12 @@ def test_reclaim_replaced_gguf_variant_prunes_old_revision_only(monkeypatch, tmp invalidated = [] monkeypatch.setattr(CI, "invalidate_hf_cache_scans", lambda: invalidated.append(True)) - result = D.reclaim_replaced_gguf_variant(repo_id, "Q4_K_M", frozenset({"NEWsha"})) + result = D.reclaim_replaced_gguf_variant( + repo_id, + "Q4_K_M", + frozenset({"NEWsha"}), + hub_cache = tmp_path, + ) assert result["removed_snapshots"] == 1 assert result["deleted_blobs"] == 1 @@ -677,13 +698,75 @@ def test_reclaim_replaced_gguf_variant_keeps_no_symlink_current_file(monkeypatch monkeypatch.setattr(CI, "all_hf_cache_scans", lambda: [SimpleNamespace(repos = [repo_info])]) monkeypatch.setattr(CI, "invalidate_hf_cache_scans", lambda: None) - result = D.reclaim_replaced_gguf_variant(repo_id, "Q4_K_M", frozenset({"REMOTEsha256"})) + result = D.reclaim_replaced_gguf_variant( + repo_id, + "Q4_K_M", + frozenset({"REMOTEsha256"}), + hub_cache = tmp_path, + ) assert snap.exists() is True # the current file must survive assert result["removed_snapshots"] == 0 assert result["deleted_blobs"] == 0 +def test_reclaim_replaced_gguf_variant_only_mutates_worker_cache(monkeypatch, tmp_path): + repo_id = "org/repo-GGUF" + cache_a = tmp_path / "cache-a" + cache_b = tmp_path / "cache-b" + + def cached_repo(cache_dir, revision): + repo_path = cache_dir / "models--org--repo-GGUF" + snap = repo_path / "snapshots" / revision / "model-Q4_K_M.gguf" + blob = repo_path / "blobs" / "OLDsha" + snap.parent.mkdir(parents = True, exist_ok = True) + blob.parent.mkdir(parents = True, exist_ok = True) + blob.write_bytes(b"old") + snap.symlink_to(blob) + return ( + SimpleNamespace( + repo_id = repo_id, + repo_type = "model", + repo_path = repo_path, + revisions = [ + SimpleNamespace( + files = [ + SimpleNamespace( + file_name = snap.name, + file_path = str(snap), + blob_path = str(blob), + ) + ] + ) + ], + ), + snap, + blob, + ) + + repo_a, snap_a, blob_a = cached_repo(cache_a, "a" * 40) + repo_b, snap_b, blob_b = cached_repo(cache_b, "b" * 40) + monkeypatch.setattr( + CI, + "all_hf_cache_scans", + lambda: [SimpleNamespace(repos = [repo_a]), SimpleNamespace(repos = [repo_b])], + ) + monkeypatch.setattr(CI, "invalidate_hf_cache_scans", lambda: None) + + result = D.reclaim_replaced_gguf_variant( + repo_id, + "Q4_K_M", + frozenset({"NEWsha"}), + hub_cache = cache_b, + ) + + assert result["removed_snapshots"] == 1 + assert snap_b.exists() is False + assert blob_b.exists() is False + assert snap_a.exists() is True + assert blob_a.exists() is True + + def _mmproj_repo(*file_names: str): return SimpleNamespace( revisions = [SimpleNamespace(files = [SimpleNamespace(file_name = n) for n in file_names])] diff --git a/studio/backend/tests/test_models_get_model_config_case_resolution.py b/studio/backend/tests/test_models_get_model_config_case_resolution.py index 12f6c497ab..a50765898b 100644 --- a/studio/backend/tests/test_models_get_model_config_case_resolution.py +++ b/studio/backend/tests/test_models_get_model_config_case_resolution.py @@ -84,7 +84,7 @@ def test_repo_in_any_hf_cache_matches_case_variant_in_legacy_cache(tmp_path, mon # covers the active cache; discard deletes case-insensitively, so detection must too, # else a decline deletes a pre-existing user repo). import utils.paths as paths_pkg - import huggingface_hub.constants as hf_constants + import hub.utils.paths as hub_paths active = tmp_path / "active" legacy = tmp_path / "legacy" @@ -96,9 +96,12 @@ def test_repo_in_any_hf_cache_matches_case_variant_in_legacy_cache(tmp_path, mon # No active-cache variant; case resolution is a no-op here. monkeypatch.setattr(paths_pkg, "resolve_cached_repo_id_case", lambda name: name) - monkeypatch.setattr(paths_pkg, "legacy_hf_cache_dir", lambda: legacy) - monkeypatch.setattr(paths_pkg, "hf_default_cache_dir", lambda: default) - monkeypatch.setattr(hf_constants, "HF_HUB_CACHE", str(active)) + monkeypatch.setattr(hub_paths, "legacy_hf_cache_dir", lambda: legacy) + monkeypatch.setattr(hub_paths, "hf_default_cache_dir", lambda: default) + monkeypatch.setattr( + "utils.hf_cache_settings.known_hf_hub_caches", + lambda: [active], + ) assert models_route._repo_in_any_hf_cache("unsloth/foo") is True # Absent from every cache -> reported absent. diff --git a/studio/backend/tests/test_offline_embedding_minimal.py b/studio/backend/tests/test_offline_embedding_minimal.py index 8862e231e5..7ff580c36d 100644 --- a/studio/backend/tests/test_offline_embedding_minimal.py +++ b/studio/backend/tests/test_offline_embedding_minimal.py @@ -85,11 +85,19 @@ def _is_embedding_model(*args, **kwargs): @pytest.fixture def hf_cache(tmp_path, monkeypatch): - """Point the HF cache at a fresh temp dir.""" + """Point the HF cache at a fresh temp dir. + + get_hf_cache_paths() reads an import-time env snapshot, not live os.environ, + so point it (and thus active_hf_hub_cache + the snapshot lookup's selected + root) at this temp cache too.""" root = tmp_path / "hub" root.mkdir() monkeypatch.setenv("HF_HOME", str(tmp_path)) monkeypatch.setenv("HF_HUB_CACHE", str(root)) + monkeypatch.setattr( + "utils.hf_cache_settings.get_hf_cache_paths", + lambda: SimpleNamespace(hub_cache = root), + ) return root @@ -198,18 +206,25 @@ def test_snapshot_dir_uses_sentence_transformers_home(tmp_path, monkeypatch): assert hf_cache_snapshot_dir("org/emb") == snapshot -def test_snapshot_dir_st_home_is_exclusive(tmp_path, monkeypatch): - # With SENTENCE_TRANSFORMERS_HOME set, ST loads only from it, so a model living only under - # HF_HUB_CACHE must not be reported. +def test_snapshot_dir_prefers_selected_cache_over_st_home(tmp_path, monkeypatch): + # The RAG loader passes cache_folder=active_hf_hub_cache(), which overrides + # SENTENCE_TRANSFORMERS_HOME, so the snapshot + offline security lookup must + # search the selected cache even when ST_HOME points elsewhere. Otherwise the + # gate scans a cache the model never loads from and a pickle weight in the + # selected cache slips through. st_home = tmp_path / "st_home" st_home.mkdir() - hub = tmp_path / "hub" - hub.mkdir() + selected = tmp_path / "hub" + selected.mkdir() monkeypatch.setenv("SENTENCE_TRANSFORMERS_HOME", str(st_home)) - monkeypatch.setenv("HF_HUB_CACHE", str(hub)) + monkeypatch.delenv("HF_HUB_CACHE", raising = False) monkeypatch.delenv("HF_HOME", raising = False) - _make_cache(hub, "org/emb", {"modules.json": MODULES_JSON}) # only in the HF hub cache - assert hf_cache_snapshot_dir("org/emb") is None + monkeypatch.setattr( + "utils.hf_cache_settings.get_hf_cache_paths", + lambda: SimpleNamespace(hub_cache = selected), + ) + snapshot = _make_cache(selected, "org/emb", {"modules.json": MODULES_JSON}) # only in selected + assert hf_cache_snapshot_dir("org/emb") == snapshot def test_snapshot_is_loadable_with_config_and_weights(hf_cache): diff --git a/studio/backend/tests/test_offline_gguf_cache_fallback.py b/studio/backend/tests/test_offline_gguf_cache_fallback.py index 35dd3979b1..d1e61d0546 100644 --- a/studio/backend/tests/test_offline_gguf_cache_fallback.py +++ b/studio/backend/tests/test_offline_gguf_cache_fallback.py @@ -130,6 +130,10 @@ def _symlink_or_skip(link: Path, target: Path) -> None: def hf_cache(tmp_path, monkeypatch): """Point ``huggingface_hub.constants.HF_HUB_CACHE`` at a temp dir.""" 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), + ) return tmp_path @@ -227,6 +231,10 @@ class TestGgufVariantFileResolution: return f"/fake/{repo_id}/{filename}" 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), + ) with ( patch( "huggingface_hub.list_repo_files", @@ -434,6 +442,40 @@ class TestGgufVariantFileResolution: assert out == str(snap / "mmproj-F16.gguf") + def test_download_companion_uses_selected_cache_not_import_time_default( + self, monkeypatch, tmp_path + ): + monkeypatch.setenv("HF_HUB_OFFLINE", "1") + import_time_cache = tmp_path / "import-time-cache" + selected_cache = tmp_path / "selected-cache" + monkeypatch.setattr(hf_constants, "HF_HUB_CACHE", str(import_time_cache)) + monkeypatch.setattr( + "utils.hf_cache_settings.get_hf_cache_paths", + lambda: _types.SimpleNamespace(hub_cache = selected_cache), + ) + repo = "unsloth/vision-GGUF" + snap = _build_cache(selected_cache, repo, {"mmproj-F16.gguf": 4}) + backend = LlamaCppBackend() + + offline_error = type("OfflineModeIsEnabled", (Exception,), {}) + + def fail_list(*_args, **_kwargs): + raise offline_error("offline") + + def fail_download(*_args, **_kwargs): + raise AssertionError("selected-cache companion must not download") + + with ( + patch("huggingface_hub.list_repo_files", fail_list), + patch( + "core.inference.llama_cpp.hf_hub_download_with_xet_fallback", + fail_download, + ), + ): + out = backend._download_mmproj(hf_repo = repo) + + assert out == str(snap / "mmproj-F16.gguf") + def test_download_includes_uppercase_split_gguf_shards(self, monkeypatch, tmp_path): backend = LlamaCppBackend() downloaded: list[str] = [] @@ -460,6 +502,10 @@ class TestGgufVariantFileResolution: return f"/fake/{repo_id}/{filename}" 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), + ) with ( patch("huggingface_hub.list_repo_files", lambda *_a, **_k: files), patch("huggingface_hub.get_paths_info", fake_get_paths_info), diff --git a/studio/backend/tests/test_openai_auto_switch.py b/studio/backend/tests/test_openai_auto_switch.py index 9361db66bb..9c6c20e6b6 100644 --- a/studio/backend/tests/test_openai_auto_switch.py +++ b/studio/backend/tests/test_openai_auto_switch.py @@ -1096,6 +1096,7 @@ def test_build_index_covers_legacy_default_lmstudio_and_custom_roots(monkeypatch from pathlib import Path import routes.models as models_route from utils import paths as upaths + from utils import hf_cache_settings import storage.studio_db as studio_db scanned = [] @@ -1116,13 +1117,18 @@ def test_build_index_covers_legacy_default_lmstudio_and_custom_roots(monkeypatch ) monkeypatch.setattr(models_route, "_resolve_hf_cache_dir", lambda: tmp_path / "active") monkeypatch.setattr(models_route, "_is_hidden_model", lambda *a, **k: False) + monkeypatch.setattr( + hf_cache_settings, + "known_hf_hub_caches", + lambda: [tmp_path / "active", tmp_path / "previous"], + ) monkeypatch.setattr(upaths, "legacy_hf_cache_dir", lambda: tmp_path / "legacy") monkeypatch.setattr(upaths, "hf_default_cache_dir", lambda: tmp_path / "default") monkeypatch.setattr(upaths, "lmstudio_model_dirs", lambda: [tmp_path / "lmstudio"]) monkeypatch.setattr( studio_db, "list_scan_folders", lambda: [{"path": str(tmp_path / "custom")}] ) - for sub in ("active", "legacy", "default", "lmstudio", "custom"): + for sub in ("active", "previous", "legacy", "default", "lmstudio", "custom"): (tmp_path / sub).mkdir() resolver._build_index() @@ -1131,6 +1137,7 @@ def test_build_index_covers_legacy_default_lmstudio_and_custom_roots(monkeypatch lm = {p for k, p in scanned if k == "lm"} assert str((tmp_path / "legacy").resolve()) in hf assert str((tmp_path / "default").resolve()) in hf + assert str((tmp_path / "previous").resolve()) in hf assert str((tmp_path / "custom").resolve()) in hf assert str((tmp_path / "lmstudio").resolve()) in lm diff --git a/studio/backend/tests/test_picker_service.py b/studio/backend/tests/test_picker_service.py index be7ea18f03..1bdfc135e3 100644 --- a/studio/backend/tests/test_picker_service.py +++ b/studio/backend/tests/test_picker_service.py @@ -241,11 +241,15 @@ def test_remote_oversized_jinja_falls_through_to_tokenizer_template(tmp_path, mo "chat_template.jinja": big_jinja, "tokenizer_config.json": tokenizer_config, } + selected_cache = tmp_path / "selected-cache" / "hub" + observed_cache_dirs = [] monkeypatch.setattr("picker.service.resolve_cached_repo_id_case", lambda name: name) monkeypatch.setattr("picker.service.iter_hf_cache_snapshots", lambda resolved: []) + monkeypatch.setattr("picker.service.active_hf_hub_cache", lambda: str(selected_cache)) def _fake_download(repo_id, rel, **kwargs): + observed_cache_dirs.append(kwargs.get("cache_dir")) target = files.get(rel) if target is None: raise FileNotFoundError(rel) @@ -264,3 +268,5 @@ def test_remote_oversized_jinja_falls_through_to_tokenizer_template(tmp_path, mo monkeypatch.setattr(huggingface_hub.HfApi, "get_paths_info", _fake_get_paths_info) assert read_default_chat_template("org/big-jinja-model") == "SMALL_TEMPLATE" + assert observed_cache_dirs + assert set(observed_cache_dirs) == {str(selected_cache)} diff --git a/studio/backend/tests/test_rag_embeddings.py b/studio/backend/tests/test_rag_embeddings.py index 28a2f69426..197ae4c495 100644 --- a/studio/backend/tests/test_rag_embeddings.py +++ b/studio/backend/tests/test_rag_embeddings.py @@ -5,8 +5,10 @@ and token counting must be serialized (else threads panic "Already borrowed").""" import os +import sys import threading import time +from types import SimpleNamespace import numpy as np import pytest @@ -130,6 +132,35 @@ def test_token_counter_enables_parallelism_only_during_call(monkeypatch): assert os.environ.get("TOKENIZERS_PARALLELISM") == "false" # restored after +def test_sentence_transformer_load_uses_live_cache(monkeypatch, tmp_path): + observed = {} + + class FakeSentenceTransformer: + def __init__(self, name, **kwargs): + observed["name"] = name + observed.update(kwargs) + + monkeypatch.setitem( + sys.modules, + "sentence_transformers", + SimpleNamespace(SentenceTransformer = FakeSentenceTransformer), + ) + monkeypatch.setattr(embeddings, "_install_torchao_stub_once", lambda: None) + monkeypatch.setattr(embeddings, "_guard_model_security", lambda *_a, **_k: None) + monkeypatch.setattr(embeddings, "_device", lambda: "cpu") + monkeypatch.setattr( + "utils.hf_cache_settings.active_hf_hub_cache", + lambda: str(tmp_path / "selected-hub"), + ) + embeddings._model = None + embeddings._name = None + + embeddings._get("Org/Embedder") + + assert observed["name"] == "Org/Embedder" + assert observed["cache_folder"] == str(tmp_path / "selected-hub") + + class _SentinelLlamaBackend: """Stand-in for LlamaServerBackend; never spawns a real server.""" diff --git a/studio/backend/tests/test_resolve_quant_gguf.py b/studio/backend/tests/test_resolve_quant_gguf.py index 840c4d8d4c..a137237e80 100644 --- a/studio/backend/tests/test_resolve_quant_gguf.py +++ b/studio/backend/tests/test_resolve_quant_gguf.py @@ -68,8 +68,6 @@ def test_skips_mtp_drafter_for_main_weights(tmp_path): def test_prefers_the_complete_snapshot(tmp_path, monkeypatch): - from huggingface_hub import constants as hf_constants - cache = tmp_path / "hub" snaps = cache / "models--org--repo" / "snapshots" # Partial older snapshot: one small shard. @@ -78,7 +76,10 @@ def test_prefers_the_complete_snapshot(tmp_path, monkeypatch): complete_first = _write(snaps / "bbbb" / "model-00001-of-00002-Q4_K_M.gguf", 30) _write(snaps / "bbbb" / "model-00002-of-00002-Q4_K_M.gguf", 40) - monkeypatch.setattr(hf_constants, "HF_HUB_CACHE", str(cache)) + monkeypatch.setattr( + "utils.hf_cache_settings.known_hf_hub_caches", + lambda: [cache], + ) path, total = models_route._resolve_quant_gguf("org/repo", "Q4_K_M", is_local = False) diff --git a/studio/backend/tests/test_setup_cache_env_hf_home.py b/studio/backend/tests/test_setup_cache_env_hf_home.py index 4520c93a51..c5d5f820a8 100644 --- a/studio/backend/tests/test_setup_cache_env_hf_home.py +++ b/studio/backend/tests/test_setup_cache_env_hf_home.py @@ -28,6 +28,9 @@ def _isolate_studio_home(monkeypatch, tmp_path): def _load_storage_roots(): + # Each test models a fresh backend process. The cache resolver intentionally + # snapshots explicit environment variables once per process. + sys.modules.pop("utils.hf_cache_settings", None) spec = importlib.util.spec_from_file_location("storage_roots_under_test", _STORAGE_ROOTS_PATH) module = importlib.util.module_from_spec(spec) spec.loader.exec_module(module) @@ -40,10 +43,10 @@ def _clear_hf_env(monkeypatch): def test_custom_hf_home_seeds_hub_and_xet(monkeypatch, tmp_path): - sr = _load_storage_roots() _clear_hf_env(monkeypatch) custom = tmp_path / "shared" / "huggingface" monkeypatch.setenv("HF_HOME", str(custom)) + sr = _load_storage_roots() sr._setup_cache_env() @@ -54,9 +57,9 @@ def test_custom_hf_home_seeds_hub_and_xet(monkeypatch, tmp_path): def test_default_when_hf_home_unset(monkeypatch, tmp_path): - sr = _load_storage_roots() _clear_hf_env(monkeypatch) monkeypatch.setenv("XDG_CACHE_HOME", str(tmp_path / "xdg")) + sr = _load_storage_roots() sr._setup_cache_env() @@ -67,11 +70,11 @@ def test_default_when_hf_home_unset(monkeypatch, tmp_path): def test_explicit_hub_cache_is_not_overridden(monkeypatch, tmp_path): - sr = _load_storage_roots() _clear_hf_env(monkeypatch) monkeypatch.setenv("HF_HOME", str(tmp_path / "home")) explicit = tmp_path / "explicit" / "hub" monkeypatch.setenv("HF_HUB_CACHE", str(explicit)) + sr = _load_storage_roots() sr._setup_cache_env() @@ -81,11 +84,11 @@ def test_explicit_hub_cache_is_not_overridden(monkeypatch, tmp_path): def test_legacy_huggingface_hub_cache_alias_is_honored(monkeypatch, tmp_path): - sr = _load_storage_roots() _clear_hf_env(monkeypatch) monkeypatch.setenv("HF_HOME", str(tmp_path / "home")) legacy = tmp_path / "legacy" / "hub" monkeypatch.setenv("HUGGINGFACE_HUB_CACHE", str(legacy)) + sr = _load_storage_roots() sr._setup_cache_env() @@ -96,15 +99,16 @@ def test_legacy_huggingface_hub_cache_alias_is_honored(monkeypatch, tmp_path): def test_whitespace_hf_home_falls_back_to_default(monkeypatch, tmp_path): # A blank/whitespace HF_HOME must not become " /hub"; fall back to default. - sr = _load_storage_roots() _clear_hf_env(monkeypatch) monkeypatch.setenv("HF_HOME", " ") monkeypatch.setenv("XDG_CACHE_HOME", str(tmp_path / "xdg")) + sr = _load_storage_roots() sr._setup_cache_env() import os + assert os.environ["HF_HOME"] == str(tmp_path / "xdg" / "huggingface") assert os.environ["HF_HUB_CACHE"] == str(tmp_path / "xdg" / "huggingface" / "hub") @@ -114,9 +118,9 @@ def test_unwritable_hf_home_does_not_crash(monkeypatch, tmp_path): blocker = tmp_path / "blocker" blocker.write_text("not a dir") unwritable = blocker / "hf" - sr = _load_storage_roots() _clear_hf_env(monkeypatch) monkeypatch.setenv("HF_HOME", str(unwritable)) + sr = _load_storage_roots() sr._setup_cache_env() # must not raise diff --git a/studio/backend/tests/test_trained_model_scan.py b/studio/backend/tests/test_trained_model_scan.py index 5d74bb7d28..64228cec3c 100644 --- a/studio/backend/tests/test_trained_model_scan.py +++ b/studio/backend/tests/test_trained_model_scan.py @@ -97,6 +97,7 @@ def test_lora_identifier_resolves_remote_adapter_base(tmp_path: Path): repo, fn, token = None, + cache_dir = None, ): assert repo == "someone/my-remote-lora" assert fn == "adapter_config.json" @@ -128,6 +129,7 @@ def test_lora_identifier_retries_transient_then_resolves(tmp_path: Path): repo, fn, token = None, + cache_dir = None, ): calls["n"] += 1 if calls["n"] == 1: diff --git a/studio/backend/tests/test_transformers_version.py b/studio/backend/tests/test_transformers_version.py index a6e6803a5c..acb2ec449b 100644 --- a/studio/backend/tests/test_transformers_version.py +++ b/studio/backend/tests/test_transformers_version.py @@ -160,6 +160,25 @@ class TestResolveBaseModel: class TestRemoteLoraBase: """_remote_lora_base reads a remote adapter's base from its Hub adapter_config.json.""" + @pytest.fixture(autouse = True) + def _selected_cache_follows_env(self, monkeypatch): + # The cache helpers now read the selected cache (get_hf_cache_paths), + # which snapshots env at import; make it follow the HF_HUB_CACHE these + # tests set so they keep driving the lookup via env. + monkeypatch.setattr( + "utils.transformers_version.get_hf_cache_paths", + lambda: _types.SimpleNamespace( + hub_cache = Path( + os.environ.get("HF_HUB_CACHE") + or os.environ.get("HUGGINGFACE_HUB_CACHE") + or os.path.join( + os.environ.get("HF_HOME") or os.path.expanduser("~/.cache/huggingface"), + "hub", + ) + ) + ), + ) + @staticmethod def _resp(cfg: dict): class _Resp: @@ -645,6 +664,24 @@ def _hf_response(cfg: dict): class TestConfigJsonHfCacheFallback: """HF hub cache is consulted only offline or after a failed fetch (never stale online).""" + @pytest.fixture(autouse = True) + def _selected_cache_follows_env(self, monkeypatch): + # As above: route the selected-cache lookup through the HF_HUB_CACHE env + # these tests set, since get_hf_cache_paths snapshots env at import. + monkeypatch.setattr( + "utils.transformers_version.get_hf_cache_paths", + lambda: _types.SimpleNamespace( + hub_cache = Path( + os.environ.get("HF_HUB_CACHE") + or os.environ.get("HUGGINGFACE_HUB_CACHE") + or os.path.join( + os.environ.get("HF_HOME") or os.path.expanduser("~/.cache/huggingface"), + "hub", + ) + ) + ), + ) + def setup_method(self): _config_json_cache.clear() diff --git a/studio/backend/tests/test_windows_external_drive_paths.py b/studio/backend/tests/test_windows_external_drive_paths.py index 9686d45c9f..5687612916 100644 --- a/studio/backend/tests/test_windows_external_drive_paths.py +++ b/studio/backend/tests/test_windows_external_drive_paths.py @@ -57,6 +57,18 @@ def test_windows_drive_roots_empty_off_windows(monkeypatch): assert external_media.windows_drive_roots() == [] +def test_macos_volume_roots_lists_readable_mounts(monkeypatch, tmp_path): + volumes = tmp_path / "Volumes" + external = volumes / "External SSD" + unreadable = volumes / "Unavailable" + external.mkdir(parents = True) + unreadable.mkdir() + monkeypatch.setattr(external_media.platform, "system", lambda: "Darwin") + monkeypatch.setattr(external_media.os, "access", lambda path, _mode: Path(path) == external) + + assert external_media.macos_volume_roots(volumes) == [external] + + def test_windows_drive_roots_lists_readable_drives(monkeypatch): _stub_windows(monkeypatch, {"C", "D", "E"}) @@ -204,8 +216,10 @@ def test_browse_allowlist_includes_windows_drive_roots(monkeypatch, tmp_path): ) fake_external_media = SimpleNamespace( linux_run_media_mount_roots = lambda: [], + macos_volume_roots = lambda: [], windows_drive_roots = lambda: [drive_root], ) + fake_paths.external_media = fake_external_media fake_studio_db = SimpleNamespace( list_scan_folders = lambda: [], contains_sensitive_path_component = lambda _p: False, @@ -270,8 +284,10 @@ def test_build_browse_allowlist_reuses_passed_roots(monkeypatch, tmp_path): ) fake_external_media = SimpleNamespace( linux_run_media_mount_roots = _media_roots, + macos_volume_roots = lambda: [], windows_drive_roots = _drive_roots, ) + fake_paths.external_media = fake_external_media fake_studio_db = SimpleNamespace(list_scan_folders = lambda: []) monkeypatch.setitem(sys.modules, "utils.paths", fake_paths) monkeypatch.setitem(sys.modules, "utils.paths.external_media", fake_external_media) diff --git a/studio/backend/utils/datasets/format_conversion.py b/studio/backend/utils/datasets/format_conversion.py index 95c9a00534..9035068c01 100644 --- a/studio/backend/utils/datasets/format_conversion.py +++ b/studio/backend/utils/datasets/format_conversion.py @@ -406,10 +406,13 @@ def convert_to_vlm_format( elif _image_lookup is not None and image_data in _image_lookup: # Bare filename โ†’ resolve via HF repo lookup from huggingface_hub import hf_hub_download + from utils.hf_cache_settings import active_hf_hub_cache + local_path = hf_hub_download( dataset_name, _image_lookup[image_data], repo_type = "dataset", + cache_dir = active_hf_hub_cache(), ) image_data = Image.open(local_path).convert("RGB") else: @@ -774,10 +777,13 @@ def convert_sharegpt_with_images_to_vlm_format( return Image.open(BytesIO(f.read())).convert("RGB") elif _image_lookup is not None and image_data in _image_lookup: from huggingface_hub import hf_hub_download + from utils.hf_cache_settings import active_hf_hub_cache + local_path = hf_hub_download( dataset_name, _image_lookup[image_data], repo_type = "dataset", + cache_dir = active_hf_hub_cache(), ) return Image.open(local_path).convert("RGB") else: diff --git a/studio/backend/utils/datasets/llm_assist.py b/studio/backend/utils/datasets/llm_assist.py index f7b35e2869..c594e883a8 100644 --- a/studio/backend/utils/datasets/llm_assist.py +++ b/studio/backend/utils/datasets/llm_assist.py @@ -58,6 +58,7 @@ def precache_helper_gguf(): try: from huggingface_hub import HfApi, hf_hub_download from huggingface_hub.utils import disable_progress_bars, enable_progress_bars + from utils.hf_cache_settings import active_hf_hub_cache disable_progress_bars() logging.getLogger("huggingface_hub").setLevel(logging.WARNING) @@ -76,7 +77,11 @@ def precache_helper_gguf(): + (f" (+{len(matching) - 1} shards)" if len(matching) > 1 else "") ) for target in matching: - hf_hub_download(repo_id = repo, filename = target) + hf_hub_download( + repo_id = repo, + filename = target, + cache_dir = active_hf_hub_cache(), + ) logger.info(f"Helper GGUF cached: {len(matching)} file(s)") else: logger.warning(f"No GGUF matching variant '{variant}' in {repo}") diff --git a/studio/backend/utils/hf_cache_settings.py b/studio/backend/utils/hf_cache_settings.py new file mode 100644 index 0000000000..07d901a3d2 --- /dev/null +++ b/studio/backend/utils/hf_cache_settings.py @@ -0,0 +1,362 @@ +# SPDX-License-Identifier: AGPL-3.0-only +# Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. See /studio/LICENSE.AGPL-3.0 + +"""Live, persisted Hugging Face cache routing for Unsloth Studio. + +Hugging Face reads cache environment variables at import time. Studio therefore +owns an explicit cache snapshot for each operation instead of trying to refresh +``huggingface_hub.constants`` in the long-running API process. +""" + +from __future__ import annotations + +import os +import shutil +import tempfile +import threading +from contextlib import contextmanager +from dataclasses import dataclass +from pathlib import Path +from typing import Iterator, Literal, Mapping, Optional + + +CACHE_HOME_SETTING_KEY = "hugging_face_cache_home" +CACHE_HISTORY_SETTING_KEY = "hugging_face_cache_history" +MAX_CACHE_HISTORY = 16 + +CacheSource = Literal["default", "studio", "environment"] + +_CACHE_ENV_KEYS = ( + "HF_HOME", + "HF_HUB_CACHE", + "HUGGINGFACE_HUB_CACHE", + "HF_XET_CACHE", +) +# Imported by storage_roots._setup_cache_env before Studio seeds defaults. +_EXPLICIT_CACHE_ENV = { + key: value.strip() + for key in _CACHE_ENV_KEYS + if (value := os.environ.get(key)) is not None and value.strip() +} +_settings_lock = threading.RLock() +_spawn_env_lock = threading.RLock() + + +@dataclass(frozen = True) +class HuggingFaceCachePaths: + cache_home: Path + hub_cache: Path + xet_cache: Path + source: CacheSource + environment_variable: Optional[str] = None + + @property + def editable(self) -> bool: + return self.source != "environment" + + @property + def is_custom(self) -> bool: + return self.source == "studio" + + def child_env(self, base: Optional[Mapping[str, str]] = None) -> dict[str, str]: + env = dict(os.environ if base is None else base) + # Do not rewrite HF_HOME. It also owns HF's token path, and credentials + # must not be moved onto a removable cache volume. + env["HF_HUB_CACHE"] = str(self.hub_cache) + env["HF_XET_CACHE"] = str(self.xet_cache) + env.pop("HUGGINGFACE_HUB_CACHE", None) + return env + + +def _default_cache_home() -> Path: + xdg = (os.environ.get("XDG_CACHE_HOME") or "").strip() + return (Path(xdg).expanduser() if xdg else Path.home() / ".cache") / "huggingface" + + +def _canonical(path: Path | str) -> Path: + return Path(path).expanduser().resolve(strict = False) + + +def _environment_paths() -> Optional[HuggingFaceCachePaths]: + explicit_home = _EXPLICIT_CACHE_ENV.get("HF_HOME") + explicit_hub = _EXPLICIT_CACHE_ENV.get("HF_HUB_CACHE") or _EXPLICIT_CACHE_ENV.get( + "HUGGINGFACE_HUB_CACHE" + ) + if not explicit_home and not explicit_hub: + return None + explicit_xet = _EXPLICIT_CACHE_ENV.get("HF_XET_CACHE") + default_home = _default_cache_home() + hf_home = _canonical(explicit_home) if explicit_home else default_home + hub = _canonical(explicit_hub) if explicit_hub else hf_home / "hub" + xet = _canonical(explicit_xet) if explicit_xet else hf_home / "xet" + controlling = next( + key + for key in ("HF_HUB_CACHE", "HUGGINGFACE_HUB_CACHE", "HF_HOME") + if key in _EXPLICIT_CACHE_ENV + ) + # Settings describes model downloads, so an explicit hub path is the + # displayed/opened location even when HF_HOME points somewhere else for + # credentials or XET data. + display_home = ( + (hub.parent if explicit_hub and hub.name.lower() == "hub" else hub) + if explicit_hub + else hf_home + ) + return HuggingFaceCachePaths(display_home, hub, xet, "environment", controlling) + + +def _stored_cache_home() -> Optional[Path]: + try: + from storage.studio_db import get_app_setting + value = get_app_setting(CACHE_HOME_SETTING_KEY, None) + except Exception: + return None + if not isinstance(value, str) or not value.strip(): + return None + try: + return _canonical(value.strip()) + except (OSError, RuntimeError, ValueError): + return None + + +def get_hf_cache_paths() -> HuggingFaceCachePaths: + env_paths = _environment_paths() + if env_paths is not None: + return env_paths + stored = _stored_cache_home() + if stored is not None: + xet = _EXPLICIT_CACHE_ENV.get("HF_XET_CACHE") + return HuggingFaceCachePaths( + stored, + stored / "hub", + _canonical(xet) if xet else stored / "xet", + "studio", + ) + home = _default_cache_home() + xet = _EXPLICIT_CACHE_ENV.get("HF_XET_CACHE") + return HuggingFaceCachePaths( + home, + home / "hub", + _canonical(xet) if xet else home / "xet", + "default", + ) + + +def active_hf_hub_cache() -> str: + """Return the current hub cache as a string for library call kwargs.""" + + return str(get_hf_cache_paths().hub_cache) + + +@contextmanager +def child_environment_for_spawn(environment: Mapping[str, str]) -> Iterator[None]: + """Apply captured env before spawn imports the child entrypoint. + + Applying variables only inside the multiprocessing target can be too late + for libraries that snapshot environment variables at import. The lock keeps + this short parent-process override atomic through ``Process.start()``. + """ + + with _spawn_env_lock: + missing = object() + saved_environment: dict[str, str | object] = {} + for key, value in environment.items(): + saved_environment[key] = os.environ.get(key, missing) + os.environ[key] = value + try: + yield + finally: + for key, previous in saved_environment.items(): + if previous is missing: + os.environ.pop(key, None) + else: + os.environ[key] = str(previous) + + +def initialize_hf_cache_environment() -> HuggingFaceCachePaths: + """Seed import-time HF variables once during backend startup.""" + + paths = get_hf_cache_paths() + # Preserve an explicit HF_HOME, otherwise keep credentials at the platform + # default while routing cache bytes through the selected home. + if not os.environ.get("HF_HOME", "").strip(): + os.environ["HF_HOME"] = str(_default_cache_home()) + os.environ["HF_HUB_CACHE"] = str(paths.hub_cache) + os.environ["HF_XET_CACHE"] = str(paths.xet_cache) + if "HUGGINGFACE_HUB_CACHE" not in _EXPLICIT_CACHE_ENV: + os.environ.pop("HUGGINGFACE_HUB_CACHE", None) + for directory in (paths.hub_cache, paths.xet_cache): + try: + directory.mkdir(parents = True, exist_ok = True) + except OSError: + pass + return paths + + +def _validate_cache_home(raw_path: str) -> Path: + value = raw_path.strip() + if not value: + raise ValueError("Choose a cache folder.") + candidate = Path(value).expanduser() + if not candidate.is_absolute(): + raise ValueError("The Hugging Face cache folder must be an absolute path.") + try: + resolved = candidate.resolve(strict = False) + except (OSError, RuntimeError, ValueError) as exc: + raise ValueError("The Hugging Face cache folder is invalid.") from exc + + if resolved.parent == resolved: + raise ValueError("Choose a folder inside the filesystem or drive root.") + try: + from hub.storage.scan_folders import ( + contains_sensitive_path_component, + is_denied_system_path, + ) + except ImportError: + contains_sensitive_path_component = is_denied_system_path = None + if is_denied_system_path is not None and is_denied_system_path(str(resolved)): + raise ValueError("System folders cannot be used for model downloads.") + if contains_sensitive_path_component is not None and contains_sensitive_path_component( + str(resolved) + ): + raise ValueError("Credential or config folders cannot be used for model downloads.") + + parent = resolved.parent + if not parent.exists() or not parent.is_dir(): + raise ValueError("The parent folder does not exist.") + try: + resolved.mkdir(exist_ok = True) + if not resolved.is_dir(): + raise ValueError("The selected cache location is not a folder.") + for child in (resolved / "hub", resolved / "xet"): + child.mkdir(exist_ok = True) + with tempfile.NamedTemporaryFile(prefix = ".unsloth-write-test-", dir = child): + pass + except PermissionError as exc: + raise ValueError("Studio does not have permission to write to this folder.") from exc + except OSError as exc: + raise ValueError(f"Studio cannot use this cache folder: {exc}") from exc + return resolved + + +def _stored_history() -> list[Path]: + try: + from storage.studio_db import get_app_setting + raw = get_app_setting(CACHE_HISTORY_SETTING_KEY, []) + except Exception: + raw = [] + if not isinstance(raw, list): + return [] + out: list[Path] = [] + seen: set[str] = set() + for value in raw: + if not isinstance(value, str) or not value.strip(): + continue + try: + path = _canonical(value) + except (OSError, RuntimeError, ValueError): + continue + key = os.path.normcase(str(path)) + if key in seen: + continue + seen.add(key) + out.append(path) + return out[:MAX_CACHE_HISTORY] + + +def set_hf_cache_home(cache_home: Optional[str]) -> HuggingFaceCachePaths: + if _environment_paths() is not None: + raise RuntimeError("The Hugging Face cache location is managed by an environment variable.") + with _settings_lock: + previous = _stored_cache_home() + next_home = _validate_cache_home(cache_home) if cache_home is not None else None + history = _stored_history() + if previous is not None and previous != next_home: + history.insert(0, previous) + deduped: list[str] = [] + seen: set[str] = set() + for path in history: + key = os.path.normcase(str(path)) + if key in seen or path == next_home: + continue + seen.add(key) + deduped.append(str(path)) + if len(deduped) >= MAX_CACHE_HISTORY: + break + from storage.studio_db import upsert_app_settings + + upsert_app_settings( + { + CACHE_HOME_SETTING_KEY: str(next_home) if next_home is not None else None, + CACHE_HISTORY_SETTING_KEY: deduped, + } + ) + # Inventory scans are cached independently from settings. Invalidate after + # persistence so the next request sees both the new active root and history. + from hub.utils.inventory_scan import invalidate_hf_cache_scans + + invalidate_hf_cache_scans() + return get_hf_cache_paths() + + +def known_hf_cache_homes() -> list[Path]: + paths = get_hf_cache_paths() + stored = _stored_cache_home() + candidates: list[Path] = [] + if paths.source != "environment": + candidates.append(paths.cache_home) + elif explicit_home := _EXPLICIT_CACHE_ENV.get("HF_HOME"): + candidates.append(_canonical(explicit_home)) + if stored is not None: + candidates.append(stored) + candidates.extend([*_stored_history(), _default_cache_home()]) + out: list[Path] = [] + seen: set[str] = set() + for candidate in candidates: + try: + canonical = _canonical(candidate) + except (OSError, RuntimeError, ValueError): + continue + key = os.path.normcase(str(canonical)) + if key in seen: + continue + seen.add(key) + out.append(canonical) + return out + + +def known_hf_hub_caches() -> list[Path]: + active = get_hf_cache_paths() + out = [active.hub_cache] + seen = {os.path.normcase(str(_canonical(active.hub_cache)))} + for home in known_hf_cache_homes(): + hub = _canonical(home / "hub") + key = os.path.normcase(str(hub)) + if key not in seen: + seen.add(key) + out.append(hub) + return out + + +def cache_status(paths: Optional[HuggingFaceCachePaths] = None) -> dict: + paths = paths or get_hf_cache_paths() + available = paths.cache_home.is_dir() + writable = available and os.access(paths.cache_home, os.W_OK | os.X_OK) + free_bytes: Optional[int] = None + if available: + try: + free_bytes = int(shutil.disk_usage(paths.cache_home).free) + except OSError: + pass + return { + "cache_home": str(paths.cache_home), + "hub_cache": str(paths.hub_cache), + "xet_cache": str(paths.xet_cache), + "source": paths.source, + "editable": paths.editable, + "is_custom": paths.is_custom, + "available": available, + "writable": writable, + "free_bytes": free_bytes, + "environment_variable": paths.environment_variable, + } diff --git a/studio/backend/utils/hf_xet_fallback.py b/studio/backend/utils/hf_xet_fallback.py index 2628b99a2d..49872f371e 100644 --- a/studio/backend/utils/hf_xet_fallback.py +++ b/studio/backend/utils/hf_xet_fallback.py @@ -21,6 +21,8 @@ never triggers the heavy load. from __future__ import annotations import threading +from functools import partial +from pathlib import Path from typing import Any, Callable, Optional # Defaults mirror unsloth_zoo.hf_xet_fallback; plain literals so they resolve (including as @@ -262,13 +264,23 @@ __all__ = [ ] -def _studio_prepare_for_http(repo_type: str, repo_id: str) -> None: +def _studio_prepare_for_http( + repo_type: str, + repo_id: str, + *, + cache_dir: Optional[str] = None, +) -> None: """Unsloth's marker-aware purge before an HTTP resume, keeping the download manager's ``.transport`` accounting consistent (vs unsloth_zoo's generic default). Guarded: a purge failure is logged, not fatal to the retry.""" try: from hub.utils.download_registry import prepare_cache_for_transport - prepare_cache_for_transport(repo_type, repo_id, "http") + prepare_cache_for_transport( + repo_type, + repo_id, + "http", + root = Path(cache_dir) if cache_dir else None, + ) except Exception as exc: try: from loggers import get_logger @@ -293,9 +305,13 @@ def hf_hub_download_with_xet_fallback( grace_period: float = DEFAULT_GRACE_PERIOD, on_status: Optional[Callable[[str], None]] = None, force_download: bool = False, + cache_dir: Optional[str] = None, ) -> str: """Single-file download via the shared fallback with Unsloth's marker-aware HTTP-retry prep. ``force_download`` re-fetches a newer blob over a cached one (Unsloth's model-update path).""" + if cache_dir is None: + from utils.hf_cache_settings import get_hf_cache_paths + cache_dir = str(get_hf_cache_paths().hub_cache) return _shared_hf_hub_download_with_xet_fallback( repo_id, filename, @@ -308,11 +324,18 @@ def hf_hub_download_with_xet_fallback( grace_period = grace_period, on_status = on_status, force_download = force_download, - prepare_for_http_fn = _studio_prepare_for_http, + cache_dir = cache_dir, + prepare_for_http_fn = partial(_studio_prepare_for_http, cache_dir = cache_dir), ) def snapshot_download_with_xet_fallback(repo_id: str, **kwargs: Any) -> str: """Whole-repo download via the shared fallback with Unsloth's marker-aware HTTP-retry prep.""" - kwargs.setdefault("prepare_for_http_fn", _studio_prepare_for_http) + if kwargs.get("cache_dir") is None: + from utils.hf_cache_settings import get_hf_cache_paths + kwargs["cache_dir"] = str(get_hf_cache_paths().hub_cache) + kwargs.setdefault( + "prepare_for_http_fn", + partial(_studio_prepare_for_http, cache_dir = kwargs["cache_dir"]), + ) return _shared_snapshot_download_with_xet_fallback(repo_id, **kwargs) diff --git a/studio/backend/utils/models/model_config.py b/studio/backend/utils/models/model_config.py index 50a997218f..4897f05ce4 100644 --- a/studio/backend/utils/models/model_config.py +++ b/studio/backend/utils/models/model_config.py @@ -37,6 +37,7 @@ import yaml from utils.native_path_leases import child_env_without_native_path_secret +from utils.hf_cache_settings import active_hf_hub_cache, get_hf_cache_paths from utils.subprocess_compat import ( windows_hidden_subprocess_kwargs as _windows_hidden_subprocess_kwargs, ) @@ -493,6 +494,7 @@ def load_model_config( trust_remote_code = trust_remote_code, token = token, local_files_only = local_files_only, + cache_dir = active_hf_hub_cache(), ) if not use_auth: @@ -503,6 +505,7 @@ def load_model_config( trust_remote_code = trust_remote_code, token = None, local_files_only = local_files_only, + cache_dir = active_hf_hub_cache(), ) # Default auth (cached tokens) @@ -510,6 +513,7 @@ def load_model_config( model_name, trust_remote_code = trust_remote_code, local_files_only = local_files_only, + cache_dir = active_hf_hub_cache(), ) @@ -624,6 +628,7 @@ def _raw_config_has_vision_config( filename = "config.json", token = hf_token, local_files_only = local_files_only, + cache_dir = active_hf_hub_cache(), ) ) config = json.loads(config_path.read_text()) @@ -770,7 +775,7 @@ def _is_vision_model_subprocess(model_name: str, hf_token: Optional[str] = None) capture_output = True, text = True, timeout = 60, - env = child_env_without_native_path_secret(), + env = get_hf_cache_paths().child_env(child_env_without_native_path_secret()), **_windows_hidden_subprocess_kwargs(), ) @@ -1714,19 +1719,20 @@ def _local_gguf_companion_search_root(selected_path: str, gguf_file: str) -> str return str(gguf_dir) -def _iter_hf_cache_snapshots(repo_id: str): +def _iter_hf_cache_snapshots(repo_id: str, cache_dir: Optional[str | Path] = None): """Yield HF cache snapshot dirs for *repo_id*, newest first. Empty if HF_HUB_CACHE is missing, the repo isn't cached, or has no snapshots. Repo name match is case-insensitive to handle casing drift between download time and lookup. """ - try: - from huggingface_hub import constants as hf_constants - except Exception: - return - - cache_dir = Path(hf_constants.HF_HUB_CACHE) + if cache_dir is None: + try: + from utils.hf_cache_settings import get_hf_cache_paths + cache_dir = get_hf_cache_paths().hub_cache + except Exception: + return + cache_dir = Path(cache_dir) target = f"models--{repo_id.replace('/', '--')}".lower() repo_dirs: list[Path] = [] try: @@ -2068,6 +2074,7 @@ def download_gguf_file( repo_id = repo_id, filename = filename, token = hf_token, + cache_dir = active_hf_hub_cache(), ) return local_path @@ -2516,7 +2523,10 @@ def get_base_model_from_lora_identifier( for _attempt in range(2): # one retry: a transient blip must not skip the base try: cfg_path = hf_hub_download( - identifier, "adapter_config.json", token = hf_token if hf_token else None + identifier, + "adapter_config.json", + token = hf_token if hf_token else None, + cache_dir = active_hf_hub_cache(), ) except (EntryNotFoundError, RepositoryNotFoundError): # No adapter_config.json -> not a resolvable LoRA; caller scans the identifier. @@ -2896,7 +2906,12 @@ class ModelConfig: try: from huggingface_hub import hf_hub_download - config_path = hf_hub_download(identifier, "adapter_config.json", token = hf_token) + config_path = hf_hub_download( + identifier, + "adapter_config.json", + token = hf_token, + cache_dir = active_hf_hub_cache(), + ) with open(config_path, "r") as f: adapter_config = json.load(f) base_model = adapter_config.get("base_model_name_or_path") diff --git a/studio/backend/utils/native_path_leases.py b/studio/backend/utils/native_path_leases.py index 08671cfe39..3ed7faa7c2 100644 --- a/studio/backend/utils/native_path_leases.py +++ b/studio/backend/utils/native_path_leases.py @@ -15,6 +15,7 @@ import base64 import binascii import hashlib import hmac +import importlib import json import os import stat as _stat_module @@ -35,7 +36,7 @@ _USED_NONCES: dict[str, int] = {} _REDACTION_LOCK = threading.Lock() _NATIVE_PATH_REDACTIONS: list[str] = [] _NATIVE_PATH_LABELS: dict[str, str] = {} -_NATIVE_PATH_ENV_LOCK = threading.Lock() +_NATIVE_PATH_ENV_LOCK = threading.RLock() _SECRET_INIT_LOCK = threading.Lock() _CACHED_LEASE_SECRET: bytes | None = None _SCRUB_REFCOUNT = 0 @@ -80,7 +81,9 @@ def child_env_without_native_path_secret(env: Mapping[str, str] | None = None) - return cleaned -def run_without_native_path_secret(target: Callable[..., Any], *args: Any, **kwargs: Any) -> Any: +def run_without_native_path_secret( + target: Callable[..., Any] | str, *args: Any, **kwargs: Any +) -> Any: """Run a multiprocessing child target without the native path lease secret.""" # Runs in the spawned child: bind it to the parent's death (Linux), since @@ -96,6 +99,11 @@ def run_without_native_path_secret(target: Callable[..., Any], *args: Any, **kwa os.environ.pop(LEASE_SECRET_ENV, None) _CACHED_LEASE_SECRET = None _SCRUB_SAVED_SECRET = None + if isinstance(target, str): + function_name, environment, *args = args + for key, value in environment.items(): + os.environ[key] = value + target = getattr(importlib.import_module(target), function_name) return target(*args, **kwargs) @@ -107,10 +115,9 @@ def native_path_secret_removed_for_child_start() -> Iterator[None]: _SCRUB_SAVED_SECRET = os.environ.pop(LEASE_SECRET_ENV, None) _CACHED_LEASE_SECRET = None _SCRUB_REFCOUNT += 1 - try: - yield - finally: - with _NATIVE_PATH_ENV_LOCK: + try: + yield + finally: _SCRUB_REFCOUNT -= 1 if _SCRUB_REFCOUNT == 0 and _SCRUB_SAVED_SECRET is not None: os.environ[LEASE_SECRET_ENV] = _SCRUB_SAVED_SECRET diff --git a/studio/backend/utils/paths/external_media.py b/studio/backend/utils/paths/external_media.py index 0ea0477cc7..1a0d2d2746 100644 --- a/studio/backend/utils/paths/external_media.py +++ b/studio/backend/utils/paths/external_media.py @@ -131,6 +131,29 @@ def linux_run_media_mount_roots( return roots +def macos_volume_roots(base: Path | str = "/Volumes") -> list[Path]: + """Readable mounted volumes for the macOS folder browser.""" + + if platform.system() != "Darwin": + return [] + base_path = Path(base) + try: + entries = list(base_path.iterdir()) + except OSError: + return [] + roots: list[Path] = [] + for entry in entries: + if is_sensitive_path_component(entry.name): + continue + try: + resolved = entry.resolve() + if resolved.is_dir() and os.access(resolved, os.R_OK | os.X_OK): + roots.append(resolved) + except (OSError, RuntimeError, ValueError): + continue + return roots + + def _active_windows_drive_bitmask() -> int: """Active-logical-drive bitmask from ``GetLogicalDrives`` (bit 0 = ``A:``), or ``0`` when unavailable. diff --git a/studio/backend/utils/paths/path_utils.py b/studio/backend/utils/paths/path_utils.py index e8dabc8954..65541661f1 100644 --- a/studio/backend/utils/paths/path_utils.py +++ b/studio/backend/utils/paths/path_utils.py @@ -122,15 +122,8 @@ def is_model_cached(model_name: str) -> bool: def _hf_hub_cache_dir() -> Path: """Return HF cache root honoring HF_HUB_CACHE when available.""" - try: - from huggingface_hub.constants import HF_HUB_CACHE - return Path(HF_HUB_CACHE) - except Exception as exc: - logger.debug( - "Could not read huggingface_hub HF_HUB_CACHE, using default hub path: %s", - exc, - ) - return Path.home() / ".cache" / "huggingface" / "hub" + from utils.hf_cache_settings import get_hf_cache_paths + return get_hf_cache_paths().hub_cache def resolve_cached_repo_id_case(model_name: str, use_memo: bool = True) -> str: diff --git a/studio/backend/utils/paths/storage_roots.py b/studio/backend/utils/paths/storage_roots.py index 35b8c57e9b..cea3cc61e3 100644 --- a/studio/backend/utils/paths/storage_roots.py +++ b/studio/backend/utils/paths/storage_roots.py @@ -277,27 +277,15 @@ def well_known_model_dirs() -> list[Path]: def _setup_cache_env() -> None: """Set cache env vars for HuggingFace, uv, and vLLM. - Respects the standard HF cache chain (explicit HF_HOME / HF_HUB_CACHE, - then XDG_CACHE_HOME, then ~/.cache/huggingface) and only sets vars the - user hasn't, so explicit overrides are honored. A user-set HF_HOME also - seeds HF_HUB_CACHE / HF_XET_CACHE (HF defaults them to $HF_HOME/hub and - $HF_HOME/xet); without this, models download to and load from the standard - cache even when HF_HOME points elsewhere, and both the Xet and HTTP-fallback - download paths inherit the same wrong root. + Explicit Hugging Face environment variables take precedence over Studio's + stored location. Studio seeds import-time variables once, while each later + worker receives its own captured cache location. """ root = cache_root() - xdg_cache = Path(os.environ.get("XDG_CACHE_HOME", Path.home() / ".cache")).expanduser() - # HUGGINGFACE_HUB_CACHE is HF's legacy alias for HF_HUB_CACHE; honor it. - if "HF_HUB_CACHE" not in os.environ and os.environ.get("HUGGINGFACE_HUB_CACHE"): - os.environ["HF_HUB_CACHE"] = os.environ["HUGGINGFACE_HUB_CACHE"] - # Seed the hub/xet caches from HF_HOME when set, else the platform default. - # Strip so a blank/whitespace HF_HOME falls back instead of making " /hub". - hf_home = (os.environ.get("HF_HOME") or "").strip() - hf_base = Path(hf_home).expanduser() if hf_home else xdg_cache / "huggingface" + from utils.hf_cache_settings import initialize_hf_cache_environment + + initialize_hf_cache_environment() defaults: dict[str, str] = { - "HF_HOME": str(hf_base), - "HF_HUB_CACHE": str(hf_base / "hub"), - "HF_XET_CACHE": str(hf_base / "xet"), "UV_CACHE_DIR": str(root / "uv"), "VLLM_CACHE_ROOT": str(root / "vllm"), } diff --git a/studio/backend/utils/security/consent.py b/studio/backend/utils/security/consent.py index b36131f809..fad52f21bb 100644 --- a/studio/backend/utils/security/consent.py +++ b/studio/backend/utils/security/consent.py @@ -147,11 +147,17 @@ def _load_remote_code_configs(model_name: str, hf_token: Optional[str] = None) - from huggingface_hub import hf_hub_download from huggingface_hub.utils import EntryNotFoundError + from utils.hf_cache_settings import active_hf_hub_cache configs = [] for name in _REMOTE_CODE_CONFIG_FILES: try: - p = hf_hub_download(repo_id = model_name, filename = name, token = hf_token) + p = hf_hub_download( + repo_id = model_name, + filename = name, + token = hf_token, + cache_dir = active_hf_hub_cache(), + ) except EntryNotFoundError: continue # genuine 404 -> truly absent except Exception: diff --git a/studio/backend/utils/security/file_security.py b/studio/backend/utils/security/file_security.py index 0490d38d7c..91d7ad8f0e 100644 --- a/studio/backend/utils/security/file_security.py +++ b/studio/backend/utils/security/file_security.py @@ -158,6 +158,7 @@ def _indexed_shard_paths( try: from huggingface_hub import hf_hub_download from huggingface_hub.utils import EntryNotFoundError + from utils.hf_cache_settings import active_hf_hub_cache except Exception: return None @@ -166,7 +167,12 @@ def _indexed_shard_paths( for prefix in _index_prefixes(load_subdirs): for filename in _TRANSFORMERS_INDEX_FILES: try: - index_path = hf_hub_download(model_name, prefix + filename, token = hf_token or None) + index_path = hf_hub_download( + model_name, + prefix + filename, + token = hf_token or None, + cache_dir = active_hf_hub_cache(), + ) except EntryNotFoundError: continue # definitively absent, not an error except Exception: diff --git a/studio/backend/utils/security/remote_code_scan.py b/studio/backend/utils/security/remote_code_scan.py index 797e1056f8..583dac94b6 100644 --- a/studio/backend/utils/security/remote_code_scan.py +++ b/studio/backend/utils/security/remote_code_scan.py @@ -29,6 +29,7 @@ from dataclasses import dataclass, field from typing import Optional from loggers import get_logger +from utils.hf_cache_settings import active_hf_hub_cache logger = get_logger(__name__) @@ -455,7 +456,12 @@ def repo_remote_code_files(model_name: str, hf_token: Optional[str] = None) -> d refs = set() for cfg_name in REMOTE_CODE_CONFIG_FILES: try: - cfg_path = hf_hub_download(model_name, cfg_name, token = hf_token) + cfg_path = hf_hub_download( + model_name, + cfg_name, + token = hf_token, + cache_dir = active_hf_hub_cache(), + ) except EntryNotFoundError: continue except Exception as exc: @@ -499,7 +505,12 @@ def repo_remote_code_files(model_name: str, hf_token: Optional[str] = None) -> d wanted = present_py | (own_refs & repo_file_set) for fn in sorted(wanted): try: - fp = hf_hub_download(model_name, fn, token = hf_token) + fp = hf_hub_download( + model_name, + fn, + token = hf_token, + cache_dir = active_hf_hub_cache(), + ) except Exception as exc: # A .py CONFIRMED PRESENT could not be fetched. A partial set would # fingerprint "clean" while transformers later runs this file, so fail @@ -602,7 +613,12 @@ def external_auto_map_repos(model_name: str, hf_token: Optional[str] = None) -> for cfg_name in REMOTE_CODE_CONFIG_FILES: try: - cfg_path = hf_hub_download(model_name, cfg_name, token = hf_token) + cfg_path = hf_hub_download( + model_name, + cfg_name, + token = hf_token, + cache_dir = active_hf_hub_cache(), + ) except EntryNotFoundError: continue except Exception: @@ -670,7 +686,12 @@ def _add_external_refs(files: dict, refs, hf_token, model_name: str) -> bool: wanted = present_py | set(entry_files) for fn in sorted(wanted): try: - fp = hf_hub_download(repo, fn, token = hf_token) + fp = hf_hub_download( + repo, + fn, + token = hf_token, + cache_dir = active_hf_hub_cache(), + ) except Exception as exc: logger.warning( "repo_remote_code_files(%s): external %s:%s unscannable (%s)", diff --git a/studio/backend/utils/transformers_version.py b/studio/backend/utils/transformers_version.py index 1fbcc9f46f..475a096248 100644 --- a/studio/backend/utils/transformers_version.py +++ b/studio/backend/utils/transformers_version.py @@ -44,6 +44,7 @@ import time from pathlib import Path from utils.native_path_leases import child_env_without_native_path_secret +from utils.hf_cache_settings import get_hf_cache_paths from utils.subprocess_compat import ( windows_hidden_subprocess_kwargs as _windows_hidden_subprocess_kwargs, ) @@ -516,13 +517,10 @@ def _adapter_base_from_hf_cache(model_name: str) -> str | None: """ if not _is_canonical_repo_id(model_name): return None - hub = ( - os.environ.get("HF_HUB_CACHE") - or os.environ.get("HUGGINGFACE_HUB_CACHE") - or os.path.join( - os.environ.get("HF_HOME") or os.path.expanduser("~/.cache/huggingface"), "hub" - ) - ) + # Route through the selected cache: after a no-restart /settings switch the + # process HF_HUB_CACHE env is stale, but the model loads from the selected + # cache, which get_hf_cache_paths() reflects (the DB switch). + hub = str(get_hf_cache_paths().hub_cache) repo_dir = Path(hub) / ("models--" + model_name.replace("/", "--")) candidates = [] ref_main = repo_dir / "refs" / "main" @@ -681,13 +679,10 @@ def _config_json_from_hf_cache(model_name: str) -> dict | None: # Only a canonical ``owner/repo`` Hub id maps to a cache dir; reject local paths. if not model_name or model_name.count("/") != 1 or model_name[0] in "/.~" or "\\" in model_name: return None - hub = ( - os.environ.get("HF_HUB_CACHE") - or os.environ.get("HUGGINGFACE_HUB_CACHE") - or os.path.join( - os.environ.get("HF_HOME") or os.path.expanduser("~/.cache/huggingface"), "hub" - ) - ) + # Route through the selected cache: after a no-restart /settings switch the + # process HF_HUB_CACHE env is stale, but the model loads from the selected + # cache, which get_hf_cache_paths() reflects (the DB switch). + hub = str(get_hf_cache_paths().hub_cache) repo_dir = Path(hub) / ("models--" + model_name.replace("/", "--")) candidates = [] ref_main = repo_dir / "refs" / "main" @@ -1251,7 +1246,7 @@ def _probe_autoconfig(target_dir: str, model_name: str, hf_token: str | None) -> True = parses, False = parse/version failure (escalate), None = transient (auth/network/offline/spawn) so the caller fails safe and does not cache. """ - env = child_env_without_native_path_secret() + env = get_hf_cache_paths().child_env(child_env_without_native_path_secret()) if hf_token: env["HF_TOKEN"] = hf_token # The probe relies on the implicit HF_TOKEN env (no token= arg). Clear any inherited @@ -1806,7 +1801,7 @@ def _install_to_dir(pkg: str, target_dir: str) -> bool: stdout = subprocess.PIPE, stderr = subprocess.STDOUT, text = True, - env = child_env_without_native_path_secret(), + env = get_hf_cache_paths().child_env(child_env_without_native_path_secret()), **_windows_hidden_subprocess_kwargs(), ) if result.returncode == 0: @@ -1829,7 +1824,7 @@ def _install_to_dir(pkg: str, target_dir: str) -> bool: stdout = subprocess.PIPE, stderr = subprocess.STDOUT, text = True, - env = child_env_without_native_path_secret(), + env = get_hf_cache_paths().child_env(child_env_without_native_path_secret()), **_windows_hidden_subprocess_kwargs(), ) if result.returncode != 0: @@ -2458,7 +2453,7 @@ def _ensure_venv_llmcompressor_exists() -> bool: stdout = subprocess.PIPE, stderr = subprocess.STDOUT, text = True, - env = child_env_without_native_path_secret(), + env = get_hf_cache_paths().child_env(child_env_without_native_path_secret()), **_windows_hidden_subprocess_kwargs(), ) last_out = result.stdout or "" diff --git a/studio/backend/utils/utils.py b/studio/backend/utils/utils.py index 21e11c6706..bf8348dd82 100644 --- a/studio/backend/utils/utils.py +++ b/studio/backend/utils/utils.py @@ -53,19 +53,41 @@ def _expand_path(raw: str) -> Path: def _hf_cache_roots() -> list: - """The one cache root the loader resolves to, by its own precedence (it picks ONE - cache_folder, no fall-through): SENTENCE_TRANSFORMERS_HOME, else HF_HUB_CACHE, else - HF_HOME/hub, else ~/.cache/huggingface/hub. Expanded, read from env, one-element list.""" - st_home = os.environ.get("SENTENCE_TRANSFORMERS_HOME") - if st_home: - return [_expand_path(st_home)] - hub = os.environ.get("HF_HUB_CACHE") or os.environ.get("HUGGINGFACE_HUB_CACHE") - if hub: - return [_expand_path(hub)] - hf_home = os.environ.get("HF_HOME") - if hf_home: - return [_expand_path(hf_home) / "hub"] - return [Path.home() / ".cache" / "huggingface" / "hub"] + """Cache roots to search for a model's local snapshot, most-authoritative first. + + The app's selected hub cache (set via /settings) is searched first: after a + no-restart cache switch the process env is stale, yet the loader reads the + selected cache via ``cache_folder=active_hf_hub_cache()``, so the snapshot + and offline security lookups must match where it actually loads. The env + precedence (SENTENCE_TRANSFORMERS_HOME, HF_HUB_CACHE, HF_HOME/hub, + ~/.cache/huggingface/hub) follows so a copy still in a previous cache resolves.""" + roots: list = [] + seen: set = set() + + def _add(path) -> None: + if path is None: + return + expanded = _expand_path(str(path)) + key = str(expanded) + if key not in seen: + seen.add(key) + roots.append(expanded) + + try: + from utils.hf_cache_settings import get_hf_cache_paths + _add(get_hf_cache_paths().hub_cache) + except Exception: + pass + + if st_home := os.environ.get("SENTENCE_TRANSFORMERS_HOME"): + _add(st_home) + if hub := (os.environ.get("HF_HUB_CACHE") or os.environ.get("HUGGINGFACE_HUB_CACHE")): + _add(hub) + if hf_home := os.environ.get("HF_HOME"): + _add(_expand_path(hf_home) / "hub") + if not roots: + _add(Path.home() / ".cache" / "huggingface" / "hub") + return roots def hf_cache_snapshot_dir(model_name: str) -> Optional[Path]: diff --git a/studio/frontend/src/features/chat/api/chat-adapter.ts b/studio/frontend/src/features/chat/api/chat-adapter.ts index b0127b5e40..d4861e8a3a 100644 --- a/studio/frontend/src/features/chat/api/chat-adapter.ts +++ b/studio/frontend/src/features/chat/api/chat-adapter.ts @@ -1404,6 +1404,7 @@ const GGUF_KNOWN_QUANT_RE = type AutoLoadCandidate = { id: string; + loadId?: string | null; kind: LastLocalModelKind; ggufVariant: string | null; maxSeqLength: number; @@ -1530,6 +1531,7 @@ async function autoLoadSmallestModel(): Promise<{ return false; } const currentStore = useChatRuntimeStore.getState(); + const modelPath = candidate.loadId ?? candidate.id; const { config } = resolveInitialConfig(candidate.id, candidate.ggufVariant); const effectiveMaxSeqLength = resolveLoadMaxSeqLength({ modelId: candidate.id, @@ -1582,7 +1584,7 @@ async function autoLoadSmallestModel(): Promise<{ : null; if ( !(await canAutoLoad({ - model_path: candidate.id, + model_path: modelPath, max_seq_length: fitMaxSeqLength, is_lora: false, gguf_variant: candidate.ggufVariant, @@ -1602,7 +1604,7 @@ async function autoLoadSmallestModel(): Promise<{ } loadAttempts += 1; const loadResp = await loadModel({ - model_path: candidate.id, + model_path: modelPath, hf_token: hfToken, max_seq_length: fitMaxSeqLength, load_in_4bit: true, @@ -1635,9 +1637,10 @@ async function autoLoadSmallestModel(): Promise<{ } // Self-gates on is_gguf (skips diffusion), so persists only for a real GGUF load. persistGpuMemoryModeOnLoad(loadResp, effectiveGpuMemoryMode); + const loadedModelId = loadResp.model || modelPath; useChatRuntimeStore .getState() - .setCheckpoint(candidate.id, candidate.ggufVariant ?? undefined); + .setCheckpoint(loadedModelId, candidate.ggufVariant ?? undefined); const store = useChatRuntimeStore.getState(); store.setModelRequiresTrustRemoteCode( loadResp.requires_trust_remote_code ?? false, @@ -1653,7 +1656,7 @@ async function autoLoadSmallestModel(): Promise<{ : effectiveMaxSeqLength, }); const autoModel: ChatModelSummary = { - id: candidate.id, + id: loadedModelId, name: loadResp.display_name ?? candidate.id, isVision: loadResp.is_vision ?? false, isLora: loadResp.is_lora ?? false, @@ -1662,7 +1665,7 @@ async function autoLoadSmallestModel(): Promise<{ audioType: loadResp.audio_type ?? null, hasAudioInput: loadResp.has_audio_input ?? false, }; - if (!store.models.some((m) => m.id === candidate.id)) { + if (!store.models.some((m) => m.id === loadedModelId)) { store.setModels([...store.models, autoModel]); } if (candidate.kind === "gguf") { @@ -1749,7 +1752,10 @@ async function autoLoadSmallestModel(): Promise<{ const repo = findCachedRepo(ggufRepos, lastLoaded.id); if (repo && lastLoaded.ggufVariant) { try { - const variants = await listGgufVariants(repo.repo_id); + const variants = await listGgufVariants(repo.repo_id, undefined, { + preferLocalCache: true, + localPath: repo.cache_path, + }); const variant = variants.variants.find( (entry) => entry.downloaded && @@ -1766,6 +1772,7 @@ async function autoLoadSmallestModel(): Promise<{ if ( await loadAutoLoadCandidate({ id: repo.repo_id, + loadId: repo.load_id, kind: "gguf", ggufVariant: variant.quant, maxSeqLength: 0, @@ -1794,6 +1801,7 @@ async function autoLoadSmallestModel(): Promise<{ if ( await loadAutoLoadCandidate({ id: repo.repo_id, + loadId: repo.load_id, kind: "model", ggufVariant: null, maxSeqLength: store.params.maxSeqLength, @@ -1823,7 +1831,10 @@ async function autoLoadSmallestModel(): Promise<{ for (const repo of sorted) { if (loadAttempts >= MAX_AUTO_LOAD_ATTEMPTS) break; try { - const variants = await listGgufVariants(repo.repo_id); + const variants = await listGgufVariants(repo.repo_id, undefined, { + preferLocalCache: true, + localPath: repo.cache_path, + }); const downloaded = variants.variants .filter((v) => v.downloaded && isAutoLoadableGgufVariant(v)) .sort((a, b) => a.size_bytes - b.size_bytes); @@ -1839,6 +1850,7 @@ async function autoLoadSmallestModel(): Promise<{ if ( await loadAutoLoadCandidate({ id: repo.repo_id, + loadId: repo.load_id, kind: "gguf", ggufVariant: variant.quant, maxSeqLength: 0, @@ -1873,6 +1885,7 @@ async function autoLoadSmallestModel(): Promise<{ if ( await loadAutoLoadCandidate({ id: repo.repo_id, + loadId: repo.load_id, kind: "model", ggufVariant: null, maxSeqLength: 4096, diff --git a/studio/frontend/src/features/chat/api/chat-api.ts b/studio/frontend/src/features/chat/api/chat-api.ts index de3e5e370c..ffaf099f29 100644 --- a/studio/frontend/src/features/chat/api/chat-api.ts +++ b/studio/frontend/src/features/chat/api/chat-api.ts @@ -243,6 +243,7 @@ export async function resolveToolConfirmation( export interface CachedGgufRepo { repo_id: string; + load_id?: string | null; size_bytes: number; cache_path: string; /** Epoch seconds of the newest downloaded quant; sorts Downloaded @@ -352,24 +353,28 @@ export async function listLocalModels( export async function listCachedGguf( signal?: AbortSignal, ): Promise { - const response = await authFetch("/api/models/cached-gguf", { signal }); + const response = await authFetch("/api/hub/cached-gguf", { signal }); const data = await parseJsonOrThrow<{ cached: CachedGgufRepo[] }>(response); return data.cached; } export interface CachedModelRepo { repo_id: string; + load_id?: string | null; size_bytes: number; /** Epoch seconds of the newest downloaded weight file; sorts Downloaded * newest-first. Optional for older-backend compatibility. */ last_modified?: number; + /** Owning cache dir; sent so a delete targets this copy, not the active + * cache. Optional for older-backend compatibility. */ + cache_path?: string | null; } export async function listCachedModels( hfToken?: string | null, signal?: AbortSignal, ): Promise { - const response = await authFetch("/api/models/cached-models", { + const response = await authFetch("/api/hub/cached-models", { headers: hubTokenHeader(hfToken), signal, }); @@ -920,8 +925,19 @@ export async function browseFolders( export async function listGgufVariants( repoId: string, hfToken?: string, + options?: { + preferLocalCache?: boolean; + localPath?: string | null; + }, ): Promise { const params = new URLSearchParams({ repo_id: repoId }); + if (options?.preferLocalCache) { + params.set("prefer_local_cache", "true"); + } + const localPath = options?.localPath?.trim(); + if (localPath) { + params.set("local_path", localPath); + } const response = await authFetch(`/api/models/gguf-variants?${params}`, { headers: hubTokenHeader(hfToken), }); diff --git a/studio/frontend/src/features/hub/catalog/dataset-download-section.tsx b/studio/frontend/src/features/hub/catalog/dataset-download-section.tsx index 81c75ca881..4a3e3f9ff6 100644 --- a/studio/frontend/src/features/hub/catalog/dataset-download-section.tsx +++ b/studio/frontend/src/features/hub/catalog/dataset-download-section.tsx @@ -50,7 +50,7 @@ export function DatasetDownloadSection({ const hfToken = useHfTokenStore((s) => s.token); const [deleteOpen, setDeleteOpen] = useState(false); const { deleting, runDelete } = useCardDelete({ - action: () => deleteCachedDataset(repoId), + action: () => deleteCachedDataset(repoId, cachePath ?? undefined), resourceName: "dataset", successMessage: () => `Deleted ${repoId}`, onSuccess: () => { diff --git a/studio/frontend/src/features/hub/catalog/gguf-download-card.tsx b/studio/frontend/src/features/hub/catalog/gguf-download-card.tsx index b3d5f45ded..870adcde11 100644 --- a/studio/frontend/src/features/hub/catalog/gguf-download-card.tsx +++ b/studio/frontend/src/features/hub/catalog/gguf-download-card.tsx @@ -765,7 +765,12 @@ export function GgufDownloadCard({ const { deleting, runDelete } = useDeleteConfirmAction({ action: async () => { if (!deleteTarget) return; - await deleteCachedModel(repoId, deleteTarget, hfToken || undefined); + await deleteCachedModel( + repoId, + deleteTarget, + hfToken || undefined, + cachePath ?? undefined, + ); }, successMessage: () => `Deleted ${repoId} ${deleteTargetLabel ?? deleteTarget}`, diff --git a/studio/frontend/src/features/hub/catalog/local-on-device-card.tsx b/studio/frontend/src/features/hub/catalog/local-on-device-card.tsx index 9f42b9ef34..9b508a5413 100644 --- a/studio/frontend/src/features/hub/catalog/local-on-device-card.tsx +++ b/studio/frontend/src/features/hub/catalog/local-on-device-card.tsx @@ -247,7 +247,10 @@ export function LocalOnDeviceCard({ const { deleting, runDelete } = useCardDelete({ action: async () => { if (!repoId) return; - await deleteCachedModel(repoId, undefined, hfToken || undefined); + // Delete is only offered for hf_cache rows (see canDelete), so `path` is + // the cache snapshot path: pass it so the delete targets the cache this + // card shows instead of falling back to the active cache. + await deleteCachedModel(repoId, undefined, hfToken || undefined, path); }, resourceName: "model", successMessage: () => `Deleted ${repoId}`, diff --git a/studio/frontend/src/features/hub/catalog/models-catalog-rows.tsx b/studio/frontend/src/features/hub/catalog/models-catalog-rows.tsx index 24b3fd49ef..5236c0be74 100644 --- a/studio/frontend/src/features/hub/catalog/models-catalog-rows.tsx +++ b/studio/frontend/src/features/hub/catalog/models-catalog-rows.tsx @@ -768,10 +768,19 @@ export const InventoryRow = memo(function InventoryRow({ ), successMessage: `Deleted ${cacheDeletableRepoId}`, onConfirm: async () => { + // Delete only the copy this row shows: cache rows carry the owning + // cache path, so pass it through and leave other caches untouched. + const rowCachePath = + row.kind === "cache" ? (row.cachePath ?? undefined) : undefined; if (isDataset) { - await deleteCachedDataset(cacheDeletableRepoId); + await deleteCachedDataset(cacheDeletableRepoId, rowCachePath); } else { - await deleteCachedModel(cacheDeletableRepoId); + await deleteCachedModel( + cacheDeletableRepoId, + undefined, + undefined, + rowCachePath, + ); // Deleted repos can't stay pinned: drop the repo pin and any of // its per-quant pins so stale rows don't linger up top. const { pinned, togglePinned: toggle } = diff --git a/studio/frontend/src/features/hub/catalog/safetensors-download-card.tsx b/studio/frontend/src/features/hub/catalog/safetensors-download-card.tsx index 0f219267c0..94c5175e8e 100644 --- a/studio/frontend/src/features/hub/catalog/safetensors-download-card.tsx +++ b/studio/frontend/src/features/hub/catalog/safetensors-download-card.tsx @@ -61,6 +61,7 @@ export function SafetensorsDownloadCard({ canRun = true, isActive, isLoadingThisModel, + cachePath, knownBytes, onLoad, onEject, @@ -75,7 +76,7 @@ export function SafetensorsDownloadCard({ canRun?: boolean; isActive: boolean; isLoadingThisModel: boolean; - /** Accepted for API parity; the options menu resolves the path itself. */ + /** Owning cache dir, threaded into delete so it targets this copy. */ cachePath?: string | null; knownBytes?: number | null; onLoad: (opts: { ggufVariant?: string; expectedBytes?: number }) => void; @@ -100,7 +101,8 @@ export function SafetensorsDownloadCard({ : null; const [deleteRepoOpen, setDeleteRepoOpen] = useState(false); const { deleting, runDelete } = useCardDelete({ - action: () => deleteCachedModel(repoId, undefined, hfToken || undefined), + action: () => + deleteCachedModel(repoId, undefined, hfToken || undefined, cachePath ?? undefined), resourceName: "model", successMessage: () => `Deleted ${repoId}`, onSuccess: () => { diff --git a/studio/frontend/src/features/hub/hooks/use-selected-model-view.ts b/studio/frontend/src/features/hub/hooks/use-selected-model-view.ts index 171dd88d3c..cceac99753 100644 --- a/studio/frontend/src/features/hub/hooks/use-selected-model-view.ts +++ b/studio/frontend/src/features/hub/hooks/use-selected-model-view.ts @@ -48,7 +48,12 @@ function localResource( ? "cached" : "local"; const id = - row.source === "hf_cache" && repoId && !row.partial ? repoId : row.loadId; + row.source === "hf_cache" && + row.activeCache !== false && + repoId && + !row.partial + ? repoId + : row.loadId; return { repoId, localPath: row.path, diff --git a/studio/frontend/src/features/hub/inventory/api.ts b/studio/frontend/src/features/hub/inventory/api.ts index 8c9214c9e5..d3a1abb540 100644 --- a/studio/frontend/src/features/hub/inventory/api.ts +++ b/studio/frontend/src/features/hub/inventory/api.ts @@ -88,6 +88,7 @@ export interface LocalModelInfo { capabilities?: BackendModelCapabilities | null; source: LocalSource; model_id?: string | null; + active_cache?: boolean | null; base_model?: string | null; base_model_source?: BaseModelSource | null; adapter_type?: string | null; @@ -258,11 +259,18 @@ export async function listCachedDatasets(): Promise { return data.cached; } -export async function deleteCachedDataset(repoId: string): Promise { +export async function deleteCachedDataset( + repoId: string, + cachePath?: string | null, +): Promise { + const payload: Record = { repo_id: repoId }; + if (cachePath) { + payload.cache_path = cachePath; + } const response = await authFetch("/api/hub/datasets/cached", { method: "DELETE", headers: { "Content-Type": "application/json" }, - body: JSON.stringify({ repo_id: repoId }), + body: JSON.stringify(payload), }); await throwIfNotOk(response, `Failed to delete dataset (${response.status})`); bumpInventoryVersion(); @@ -272,11 +280,17 @@ export async function deleteCachedModel( repoId: string, variant?: string, hfToken?: string | null, + cachePath?: string | null, ): Promise { const payload: Record = { repo_id: repoId }; if (variant) { payload.variant = variant; } + // Scope the delete to the exact cache this row represents so copies in other, + // previously selected caches are not removed. + if (cachePath) { + payload.cache_path = cachePath; + } const response = await authFetch("/api/hub/delete-cached", { method: "DELETE", headers: { "Content-Type": "application/json", ...hubTokenHeader(hfToken) }, diff --git a/studio/frontend/src/features/hub/inventory/types.ts b/studio/frontend/src/features/hub/inventory/types.ts index b300be5fb0..d1f2f993cd 100644 --- a/studio/frontend/src/features/hub/inventory/types.ts +++ b/studio/frontend/src/features/hub/inventory/types.ts @@ -83,6 +83,7 @@ export interface LocalInventoryRow { updatedAt: number | null; partial?: boolean; partialTransport?: string | null; + activeCache?: boolean | null; pipelineTag?: string | null; tags?: string[]; libraryName?: string | null; diff --git a/studio/frontend/src/features/hub/inventory/view-models.ts b/studio/frontend/src/features/hub/inventory/view-models.ts index 5fd9cc2228..41b4733d19 100644 --- a/studio/frontend/src/features/hub/inventory/view-models.ts +++ b/studio/frontend/src/features/hub/inventory/view-models.ts @@ -301,6 +301,7 @@ export function buildLocalInventoryRows( updatedAt: normalizeTimestamp(model.updated_at), partial: model.partial ?? false, partialTransport: model.partial_transport ?? null, + activeCache: model.active_cache ?? null, pipelineTag: model.pipeline_tag ?? null, tags: model.tags, libraryName: model.library_name ?? null, diff --git a/studio/frontend/src/features/model-picker/components/model-selector/folder-browser.tsx b/studio/frontend/src/features/model-picker/components/model-selector/folder-browser.tsx index 00a3e5d563..bad6c27bb2 100644 --- a/studio/frontend/src/features/model-picker/components/model-selector/folder-browser.tsx +++ b/studio/frontend/src/features/model-picker/components/model-selector/folder-browser.tsx @@ -28,6 +28,9 @@ export interface FolderBrowserProps { onSelect: (path: string) => void; /** Optional initial directory. Defaults to the user's home on the server. */ initialPath?: string; + title?: string; + confirmLabel?: string; + showModelHints?: boolean; } function splitBreadcrumb(path: string): { label: string; value: string }[] { @@ -78,6 +81,9 @@ export function FolderBrowser({ onOpenChange, onSelect, initialPath, + title = "Select folder to detect models", + confirmLabel = "Use this folder", + showModelHints = true, }: FolderBrowserProps) { const [data, setData] = useState(null); const [path, setPath] = useState(initialPath); @@ -150,7 +156,7 @@ export function FolderBrowser({ data-testid="folder-browser-dialog" > - Select folder to detect models + {title} {/* Breadcrumb */} @@ -230,12 +236,13 @@ export function FolderBrowser({ )} {data.entries.length === 0 && - !(data.model_files_here && data.model_files_here > 0) && ( + (!showModelHints || + !(data.model_files_here && data.model_files_here > 0)) && (
(empty directory)
)} - {data.model_files_here !== undefined && + {showModelHints && data.model_files_here !== undefined && data.model_files_here > 0 && (
{data.model_files_here} model file @@ -272,7 +279,7 @@ export function FolderBrowser({ )} /> {e.name} - {e.has_models && ( + {showModelHints && e.has_models && ( models @@ -312,7 +319,7 @@ export function FolderBrowser({ onClick={handleConfirm} disabled={!path || loading || !!error} > - Use this folder + {confirmLabel}
diff --git a/studio/frontend/src/features/model-picker/components/model-selector/pickers.tsx b/studio/frontend/src/features/model-picker/components/model-selector/pickers.tsx index aeb0fcf40f..20810b3f11 100644 --- a/studio/frontend/src/features/model-picker/components/model-selector/pickers.tsx +++ b/studio/frontend/src/features/model-picker/components/model-selector/pickers.tsx @@ -2912,7 +2912,12 @@ export function HubModelPicker({ updateGgufVariant(c.repo_id, quant, expectedBytes), updateDisabled: loadedModelId === c.repo_id, onDelete: async (quant) => { - await deleteCachedModel(c.repo_id, quant, hfToken || undefined); + await deleteCachedModel( + c.repo_id, + quant, + hfToken || undefined, + c.cache_path || undefined, + ); prunePinnedQuantValidation(c.repo_id, quant); refreshCachedLists(); }, @@ -2990,7 +2995,12 @@ export function HubModelPicker({ successMessage: `Deleted ${c.repo_id}`, disabled: deleteDisabled, onConfirm: async () => { - await deleteCachedModel(c.repo_id, undefined, hfToken || undefined); + await deleteCachedModel( + c.repo_id, + undefined, + hfToken || undefined, + c.cache_path || undefined, + ); if (pinnedSet.has(pinKey(c.repo_id))) { togglePinned(c.repo_id); } diff --git a/studio/frontend/src/features/native-intents/api.ts b/studio/frontend/src/features/native-intents/api.ts index b5c7e33b3a..98caa709b9 100644 --- a/studio/frontend/src/features/native-intents/api.ts +++ b/studio/frontend/src/features/native-intents/api.ts @@ -23,6 +23,11 @@ export async function pickNativeModel(): Promise { return invokeNative("pick_native_model"); } +export async function pickHuggingFaceCacheDir(): Promise { + if (!isTauri) return null; + return invokeNative("pick_hugging_face_cache_dir"); +} + export async function registerNativeModelPath(path: string): Promise { return invokeNative("register_native_model_path", { path }); } diff --git a/studio/frontend/src/features/native-intents/index.ts b/studio/frontend/src/features/native-intents/index.ts index a82c39a9ec..98b89f08f3 100644 --- a/studio/frontend/src/features/native-intents/index.ts +++ b/studio/frontend/src/features/native-intents/index.ts @@ -3,7 +3,7 @@ export { NativeModelChip } from "./components/native-model-chip"; export { NativeModelDropOverlay } from "./components/native-model-drop-overlay"; -export { openModelsDir } from "./api"; +export { openModelsDir, pickHuggingFaceCacheDir } from "./api"; export { useNativeIntentStore } from "./store"; export type { NativeIntent } from "./types"; export { useChooseNativeModel } from "./use-native-dialogs"; diff --git a/studio/frontend/src/features/settings/api/hugging-face-cache.ts b/studio/frontend/src/features/settings/api/hugging-face-cache.ts new file mode 100644 index 0000000000..3ca0f31053 --- /dev/null +++ b/studio/frontend/src/features/settings/api/hugging-face-cache.ts @@ -0,0 +1,90 @@ +// SPDX-License-Identifier: AGPL-3.0-only +// Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. See /studio/LICENSE.AGPL-3.0 + +import { authFetch } from "@/features/auth"; +import { + bumpInventoryVersion, + invalidateGgufVariantsCache, +} from "@/features/hub"; +import { readFastApiError } from "@/lib/format-fastapi-error"; + +export type HuggingFaceCacheSettings = { + cacheHome: string; + hubCache: string; + xetCache: string; + source: "default" | "studio" | "environment"; + editable: boolean; + isCustom: boolean; + available: boolean; + writable: boolean; + freeBytes: number | null; + environmentVariable: string | null; +}; + +type ApiHuggingFaceCacheSettings = { + // biome-ignore lint/style/useNamingConvention: API schema + cache_home: string; + // biome-ignore lint/style/useNamingConvention: API schema + hub_cache: string; + // biome-ignore lint/style/useNamingConvention: API schema + xet_cache: string; + source: HuggingFaceCacheSettings["source"]; + editable: boolean; + // biome-ignore lint/style/useNamingConvention: API schema + is_custom: boolean; + available: boolean; + writable: boolean; + // biome-ignore lint/style/useNamingConvention: API schema + free_bytes: number | null; + // biome-ignore lint/style/useNamingConvention: API schema + environment_variable: string | null; +}; + +function fromApi(value: ApiHuggingFaceCacheSettings): HuggingFaceCacheSettings { + return { + cacheHome: value.cache_home, + hubCache: value.hub_cache, + xetCache: value.xet_cache, + source: value.source, + editable: value.editable, + isCustom: value.is_custom, + available: value.available, + writable: value.writable, + freeBytes: value.free_bytes, + environmentVariable: value.environment_variable, + }; +} + +export async function loadHuggingFaceCacheSettings() { + const response = await authFetch("/api/settings/hugging-face-cache"); + if (!response.ok) { + throw new Error( + await readFastApiError( + response, + "Failed to load the model cache location", + ), + ); + } + return fromApi(await response.json()); +} + +export async function updateHuggingFaceCacheSettings(cacheHome: string | null) { + const response = await authFetch("/api/settings/hugging-face-cache", { + method: "PUT", + headers: { "Content-Type": "application/json" }, + // biome-ignore lint/style/useNamingConvention: API schema + body: JSON.stringify({ cache_home: cacheHome }), + }); + if (!response.ok) { + throw new Error( + await readFastApiError( + response, + "Failed to update the model cache location", + ), + ); + } + const settings = fromApi(await response.json()); + bumpInventoryVersion(); + invalidateGgufVariantsCache(); + return settings; +} diff --git a/studio/frontend/src/features/settings/api/models-folder.ts b/studio/frontend/src/features/settings/api/models-folder.ts deleted file mode 100644 index 13e8c065c3..0000000000 --- a/studio/frontend/src/features/settings/api/models-folder.ts +++ /dev/null @@ -1,40 +0,0 @@ -// SPDX-License-Identifier: AGPL-3.0-only -// Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. See /studio/LICENSE.AGPL-3.0 - -import { authFetch } from "@/features/auth"; -import { readFastApiError } from "@/lib/format-fastapi-error"; - -export type ModelsFolder = { - path: string; -}; - -// The path is resolved once at backend startup and never changes, so cache it -// and dedupe concurrent loads (same shape as the sibling settings loaders). -let cachedModelsFolder: ModelsFolder | null = null; -let inFlightModelsFolder: Promise | null = null; - -async function fetchModelsFolder(): Promise { - const res = await authFetch("/api/hub/models-folder"); - if (!res.ok) { - throw new Error( - await readFastApiError(res, "Failed to load models folder"), - ); - } - const data = (await res.json()) as { path: string }; - return { path: data.path }; -} - -export async function loadModelsFolder(): Promise { - if (cachedModelsFolder) { - return cachedModelsFolder; - } - inFlightModelsFolder ??= fetchModelsFolder() - .then((folder) => { - cachedModelsFolder = folder; - return folder; - }) - .finally(() => { - inFlightModelsFolder = null; - }); - return inFlightModelsFolder; -} diff --git a/studio/frontend/src/features/settings/settings-search.ts b/studio/frontend/src/features/settings/settings-search.ts index c19b27c732..380909e696 100644 --- a/studio/frontend/src/features/settings/settings-search.ts +++ b/studio/frontend/src/features/settings/settings-search.ts @@ -16,8 +16,6 @@ export const SETTINGS_SEARCH_INDEX: Record = { "settings.general.huggingFaceToken", "settings.general.gettingStarted", "settings.general.startOnboarding", - "settings.general.storage.sectionTitle", - "settings.general.storage.modelsFolder", "settings.appearance.language.title", "settings.appearance.language.label", "settings.general.notifications.sectionTitle", @@ -71,6 +69,7 @@ export const SETTINGS_SEARCH_INDEX: Record = { "settings.resources.gpu.title", "settings.resources.storage.title", "settings.resources.storage.modelsFolder", + "settings.resources.storage.futureDownloads", "settings.resources.storage.systemDisk", "settings.resources.environment.title", "settings.resources.environment.backend", diff --git a/studio/frontend/src/features/settings/tabs/general-tab.tsx b/studio/frontend/src/features/settings/tabs/general-tab.tsx index 17697f933c..1d5c533a25 100644 --- a/studio/frontend/src/features/settings/tabs/general-tab.tsx +++ b/studio/frontend/src/features/settings/tabs/general-tab.tsx @@ -15,7 +15,6 @@ import { Switch } from "@/components/ui/switch"; import { usePlatformStore } from "@/config/env"; import { resetOnboardingDone } from "@/features/auth"; import { PermissionModeDropdown, useChatRuntimeStore } from "@/features/chat"; -import { openModelsDir } from "@/features/native-intents"; import { emitTrainingRunsChanged } from "@/features/training"; import { setShowLlamaUpdateBanner, @@ -24,7 +23,6 @@ import { import { useHfTokenValidation } from "@/hooks"; import { LOCALE_STORAGE_KEY, useT } from "@/i18n"; import { isTauri } from "@/lib/api-base"; -import { copyToClipboard } from "@/lib/copy-to-clipboard"; import { toast } from "@/lib/toast"; import { cn } from "@/lib/utils"; import { useNavigate, useRouterState } from "@tanstack/react-router"; @@ -43,7 +41,6 @@ import { loadHelperPrecacheSettings, updateHelperPrecacheSettings, } from "../api/helper-precache"; -import { type ModelsFolder, loadModelsFolder } from "../api/models-folder"; import { type PreviewSharingSettings, loadPreviewSharing, @@ -183,7 +180,6 @@ export function GeneralTab() { const [isSavingPreviewSharing, setIsSavingPreviewSharing] = useState(false); const [revokePreviewOpen, setRevokePreviewOpen] = useState(false); const [isRevokingPreview, setIsRevokingPreview] = useState(false); - const [modelsFolder, setModelsFolder] = useState(null); const [embeddingModel, setEmbeddingModel] = useState(null); const [draftEmbeddingModel, setDraftEmbeddingModel] = useState(""); @@ -317,43 +313,6 @@ export function GeneralTab() { }; }, [t]); - useEffect(() => { - let cancelled = false; - void loadModelsFolder() - .then((folder) => { - if (cancelled) return; - setModelsFolder(folder); - }) - .catch(() => { - // Non-critical: leave the row hidden if the path can't be resolved. - }); - return () => { - cancelled = true; - }; - }, []); - - // Desktop opens the folder in the OS file manager; the browser can't, so it - // falls back to copying the path (which is the info users actually want). - const handleModelsFolder = async () => { - const folder = modelsFolder; - if (!folder) return; - if (isTauri) { - try { - await openModelsDir(folder.path); - } catch (error) { - toast.error(t("settings.general.storage.openError"), { - description: error instanceof Error ? error.message : undefined, - }); - } - return; - } - if (await copyToClipboard(folder.path)) { - toast.success(t("settings.general.storage.copied")); - } else { - toast.error(t("settings.general.storage.copyError")); - } - }; - const saveHelperPrecache = async (enabled: boolean) => { setIsSavingHelperPrecache(true); setHelperPrecacheError(null); @@ -588,33 +547,6 @@ export function GeneralTab() { )} - {modelsFolder ? ( - - -
- - {modelsFolder.path} - - -
-
-
- ) : null} - = 10 ? gib.toFixed(1) : gib.toFixed(2)} GiB`; +} + // RAM/VRAM come from the backend in binary units (bytes / 1024**3), matching // nvidia-smi and PyTorch, so label those readouts GiB. Disk stays on formatGb // because the backend reports disk in decimal GB (bytes / 1e9). @@ -165,20 +180,22 @@ export function ResourcesTab() { enabled: liveUpdates, pollMs: liveUpdates ? POLL_MS : undefined, }); - const [modelsFolder, setModelsFolder] = useState(null); - const [modelsFolderLoaded, setModelsFolderLoaded] = useState(false); + const [hfCache, setHfCache] = useState(null); + const [hfCacheLoaded, setHfCacheLoaded] = useState(false); + const [cacheBrowserOpen, setCacheBrowserOpen] = useState(false); + const [cacheSaving, setCacheSaving] = useState(false); useEffect(() => { let cancelled = false; - void loadModelsFolder() - .then((folder) => { + void loadHuggingFaceCacheSettings() + .then((settings) => { if (cancelled) return; - setModelsFolder(folder); - setModelsFolderLoaded(true); + setHfCache(settings); + setHfCacheLoaded(true); }) .catch(() => { if (cancelled) return; - setModelsFolderLoaded(true); + setHfCacheLoaded(true); }); return () => { cancelled = true; @@ -237,12 +254,11 @@ export function ResourcesTab() { }; }, [systemInfo]); - const handleModelsFolder = async () => { - const folder = modelsFolder; - if (!folder) return; + const handleCacheFolder = async () => { + if (!hfCache) return; if (isTauri) { try { - await openModelsDir(folder.path); + await openModelsDir(hfCache.cacheHome); } catch (error) { toast.error(t("settings.resources.storage.openError"), { description: error instanceof Error ? error.message : undefined, @@ -250,13 +266,43 @@ export function ResourcesTab() { } return; } - if (await copyToClipboard(folder.path)) { + if (await copyToClipboard(hfCache.cacheHome)) { toast.success(t("settings.resources.storage.copied")); } else { toast.error(t("settings.resources.storage.copyError")); } }; + const saveCacheFolder = async (path: string | null) => { + setCacheSaving(true); + try { + const settings = await updateHuggingFaceCacheSettings(path); + setHfCache(settings); + toast.success(t("settings.resources.storage.cacheSaved")); + } catch (error) { + toast.error(t("settings.resources.storage.cacheSaveError"), { + description: error instanceof Error ? error.message : undefined, + }); + } finally { + setCacheSaving(false); + } + }; + + const changeCacheFolder = async () => { + if (!isTauri) { + setCacheBrowserOpen(true); + return; + } + try { + const path = await pickHuggingFaceCacheDir(); + if (path) await saveCacheFolder(path); + } catch (error) { + toast.error(t("settings.resources.storage.cachePickerError"), { + description: error instanceof Error ? error.message : undefined, + }); + } + }; + const cpuCoresLabel = systemInfo.cpu?.logical_count && systemInfo.cpu?.physical_count ? t("settings.resources.liveMonitor.cpuCores", { @@ -270,11 +316,27 @@ export function ResourcesTab() { const backendLabel = ( systemInfo.gpu?.backend ?? systemInfo.device_backend ?? "cpu" ).toUpperCase(); - const modelsFolderPath = modelsFolder - ? modelsFolder.path - : modelsFolderLoaded + const modelsFolderPath = hfCache + ? hfCache.cacheHome + : hfCacheLoaded ? t("settings.resources.environment.unknown") : t("common.loading"); + const cacheLocationDetail = hfCache + ? hfCache.source === "environment" + ? t("settings.resources.storage.environmentManaged", { + variable: hfCache.environmentVariable ?? "HF_HOME", + }) + : [ + t("settings.resources.storage.futureDownloads"), + hfCache.freeBytes !== null + ? t("settings.resources.storage.locationFree", { + free: formatBytes(hfCache.freeBytes) ?? "", + }) + : null, + ] + .filter(Boolean) + .join(" ยท ") + : null; const unknownLabel = t("settings.resources.environment.unknown"); return ( @@ -467,29 +529,86 @@ export function ResourcesTab() { -
- - {modelsFolderPath} - +
+
+ + +
+ {cacheLocationDetail || hfCache?.isCustom ? ( +
+ {cacheLocationDetail ? ( + + {cacheLocationDetail} + + ) : null} + {hfCache?.isCustom ? ( + + ) : null} +
+ ) : null}
+ void saveCacheFolder(path)} + initialPath={hfCache?.cacheHome} + title={t("settings.resources.storage.chooseTitle")} + confirmLabel={t("settings.resources.storage.chooseAction")} + showModelHints={false} + /> + String { + if let Some(rest) = path.strip_prefix(r"\\?\UNC\") { + return format!(r"\\{rest}"); + } + path.strip_prefix(r"\\?\").unwrap_or(&path).to_string() +} + +fn portable_path_string(path: &Path) -> String { + let value = path.to_string_lossy().to_string(); + #[cfg(windows)] + { + return normalize_windows_verbatim_path(value); + } + #[cfg(not(windows))] + { + value + } +} + #[derive(Clone, Debug)] struct NativePathEntry { token: String, @@ -311,6 +330,34 @@ pub async fn pick_native_model( .map(Some) } +#[tauri::command] +pub async fn pick_hugging_face_cache_dir( + window: WebviewWindow, + app: AppHandle, +) -> Result, String> { + ensure_main_window(&window)?; + let (tx, rx) = tokio::sync::oneshot::channel(); + app.dialog() + .file() + .set_title("Choose model download location") + .pick_folder(move |path| { + let _ = tx.send(path); + }); + let Some(folder_path) = rx.await.map_err(|_| "Dialog closed".to_string())? else { + return Ok(None); + }; + let path = folder_path + .into_path() + .map_err(|_| "Only local filesystem folders are supported.".to_string())?; + let canonical = path + .canonicalize() + .map_err(|e| format!("Could not use the selected folder: {e}"))?; + if !canonical.is_dir() { + return Err("The selected location is not a folder.".to_string()); + } + Ok(Some(portable_path_string(&canonical))) +} + #[tauri::command] pub fn consume_native_path_token( window: WebviewWindow, @@ -451,6 +498,18 @@ mod tests { let _ = fs::remove_file(path); } + #[test] + fn windows_verbatim_paths_are_portable() { + assert_eq!( + normalize_windows_verbatim_path(r"\\?\C:\models\cache".to_string()), + r"C:\models\cache" + ); + assert_eq!( + normalize_windows_verbatim_path(r"\\?\UNC\server\share\cache".to_string()), + r"\\server\share\cache" + ); + } + #[cfg(unix)] #[test] fn reveal_rejects_symlink_replacement() { diff --git a/tests/studio/test_model_picker_contracts.py b/tests/studio/test_model_picker_contracts.py index 20335a279c..baf0fdf1bf 100644 --- a/tests/studio/test_model_picker_contracts.py +++ b/tests/studio/test_model_picker_contracts.py @@ -60,6 +60,19 @@ def test_compare_load_clears_stale_native_lease(): assert "activeNativePathExpiresAtMs: null" in src +def test_autoload_records_backend_loaded_model_identity(): + """An inactive-cache inventory row loads by local path, so startup autoload + must key both the active checkpoint and its summary by the backend's loaded + model identity instead of the catalog repo id.""" + src = _read("features/chat/api/chat-adapter.ts") + autoload = src.split("async function loadAutoLoadCandidate", 1)[1] + autoload = autoload.split("\n try {", 1)[0] + assert "const loadedModelId = loadResp.model || modelPath" in autoload + assert "setCheckpoint(loadedModelId," in autoload + assert "id: loadedModelId" in autoload + assert "m.id === loadedModelId" in autoload + + def test_rollback_restores_native_lease_expiry_with_token(): """A failed model switch that rolls back to a previously loaded picked GGUF must restore the lease expiry paired with the token, never the token alone @@ -248,6 +261,30 @@ def test_pinned_validation_uses_cached_local_variant_listing(): assert "bumpInventoryVersion(" in delete_fn +def test_chat_autoload_scopes_variant_lookup_to_cached_repo_path(): + """Autoload must probe the exact cache row it will load, including rows + retained from a previously selected Hugging Face cache.""" + src = _read("features/chat/api/chat-adapter.ts") + auto_load = src.split("async function autoLoadSmallestModel", 1)[1] + assert auto_load.count("preferLocalCache: true") >= 2 + assert auto_load.count("localPath: repo.cache_path") >= 2 + + chat_api = _read("features/chat/api/chat-api.ts") + variants_fn = chat_api.split("export async function listGgufVariants", 1)[1] + variants_fn = variants_fn.split("export interface KvCacheEstimate", 1)[0] + assert 'params.set("prefer_local_cache", "true")' in variants_fn + assert 'params.set("local_path", localPath)' in variants_fn + + +def test_cache_location_update_invalidates_frontend_inventory(): + """A successful cache switch must refresh both inventory rows and cached + GGUF variant results before any stale active-cache identity can be reused.""" + src = _read("features/settings/api/hugging-face-cache.ts") + update_fn = src.split("export async function updateHuggingFaceCacheSettings", 1)[1] + assert "bumpInventoryVersion();" in update_fn + assert "invalidateGgufVariantsCache();" in update_fn + + def test_downloaded_list_offsets_virtual_rows(): """The On Device virtualized list sits below the Pinned block in the same scroll element, so it must pass its measured offset as scrollMargin or rows