From a84931666dc7a9d6477114676e34c35d879e1d8b Mon Sep 17 00:00:00 2001 From: Snakinya <70094531+Snakinya@users.noreply.github.com> Date: Mon, 21 Sep 2026 20:04:23 +0800 Subject: [PATCH] [Fix] Scope artifact caches to registry endpoint Keep the historical cache layout for the default service while assigning every custom registry a stable endpoint namespace. Apply the same namespace to downloads, offline lookup, verification, scanning, clearing, and download locks. Add regression coverage for cross-registry artifact isolation, legacy-cache handling, cache management, and verification. --- src/modelscope_hub/_cache_manager.py | 19 +++-- src/modelscope_hub/_cache_paths.py | 24 ++++++ src/modelscope_hub/_download.py | 9 ++- src/modelscope_hub/api.py | 8 +- tests/test_endpoint_cache_namespace.py | 100 +++++++++++++++++++++++++ 5 files changed, 151 insertions(+), 9 deletions(-) create mode 100644 src/modelscope_hub/_cache_paths.py create mode 100644 tests/test_endpoint_cache_namespace.py diff --git a/src/modelscope_hub/_cache_manager.py b/src/modelscope_hub/_cache_manager.py index 980ae75..167a1c1 100644 --- a/src/modelscope_hub/_cache_manager.py +++ b/src/modelscope_hub/_cache_manager.py @@ -9,8 +9,9 @@ import shutil from pathlib import Path +from ._cache_paths import endpoint_cache_root from .config import get_default_config -from .constants import RepoType +from .constants import DEFAULT_ENDPOINT, RepoType from .errors import CacheError from .types import CachedRepoInfo, CacheInfo, CacheVerification, VerificationMismatch from .utils.file_utils import compute_hash @@ -22,7 +23,7 @@ _DEFAULT_SCAN_TYPES = [RepoType.MODEL, RepoType.DATASET, RepoType.STUDIO, RepoType.MCP] -def scan_cache(cache_dir: Path | None = None) -> CacheInfo: +def scan_cache(cache_dir: Path | None = None, *, endpoint: str | None = DEFAULT_ENDPOINT) -> CacheInfo: """Scan the local cache and return metadata about cached repositories. Parameters @@ -36,7 +37,7 @@ def scan_cache(cache_dir: Path | None = None) -> CacheInfo: Summary of all cached repositories, total size, etc. """ config = get_default_config() - root = Path(cache_dir) if cache_dir else config.cache_dir + root = endpoint_cache_root(Path(cache_dir) if cache_dir else config.cache_dir, endpoint) if not root.is_dir(): return CacheInfo(repos=[], total_size=0, cache_dir=str(root)) @@ -136,6 +137,8 @@ def clear_cache( cache_dir: Path | None = None, repo_type: str | None = None, repo_id: str | None = None, + *, + endpoint: str | None = DEFAULT_ENDPOINT, ) -> int: """Remove cached data from disk. @@ -160,7 +163,7 @@ def clear_cache( On filesystem errors. """ config = get_default_config() - root = Path(cache_dir) if cache_dir else config.cache_dir + root = endpoint_cache_root(Path(cache_dir) if cache_dir else config.cache_dir, endpoint) # Guard against accidental nuke: passing only ``repo_id`` would otherwise # silently fall through to the "clear everything" branch below. @@ -223,6 +226,7 @@ def verify_cache( revision: str | None = None, cache_dir: str | Path | None = None, local_dir: str | Path | None = None, + endpoint: str | None = DEFAULT_ENDPOINT, ) -> CacheVerification: """Compare a cached snapshot or local directory with remote SHA-256 values.""" root, resolved_revision = _resolve_verification_root( @@ -231,6 +235,7 @@ def verify_cache( revision=revision, cache_dir=cache_dir, local_dir=local_dir, + endpoint=endpoint, ) local_by_path = { relative.as_posix(): path @@ -275,6 +280,7 @@ def _resolve_verification_root( revision: str | None, cache_dir: str | Path | None, local_dir: str | Path | None, + endpoint: str | None = DEFAULT_ENDPOINT, ) -> tuple[Path, str]: if local_dir is not None: root = Path(local_dir).expanduser().resolve() @@ -282,7 +288,10 @@ def _resolve_verification_root( raise CacheError(f"Local directory does not exist: {root}") return root, revision or "master" - cache_root = Path(cache_dir or get_default_config().cache_dir).expanduser().resolve() + cache_root = endpoint_cache_root( + Path(cache_dir or get_default_config().cache_dir).expanduser().resolve(), + endpoint, + ) segment = f"{repo_type}s" if not repo_type.endswith("s") else repo_type repo_root = cache_root / segment / repo_id.replace("/", "--") snapshots = repo_root / "snapshots" diff --git a/src/modelscope_hub/_cache_paths.py b/src/modelscope_hub/_cache_paths.py new file mode 100644 index 0000000..0b38121 --- /dev/null +++ b/src/modelscope_hub/_cache_paths.py @@ -0,0 +1,24 @@ +"""Shared cache-path helpers.""" + +from __future__ import annotations + +import hashlib +from pathlib import Path + +from .config import HubConfig +from .constants import DEFAULT_ENDPOINT + + +def endpoint_cache_root(cache_dir: str | Path, endpoint: str | None) -> Path: + """Return the cache root reserved for *endpoint*. + + Keep the historical layout for the default service so existing caches stay + usable. Other registries receive an opaque namespace derived from their + normalized endpoint and therefore cannot reuse each other's artifacts. + """ + root = Path(cache_dir) + normalized = HubConfig.normalize_endpoint(endpoint) + if normalized.casefold() == DEFAULT_ENDPOINT.casefold(): + return root + namespace = hashlib.sha256(normalized.encode("utf-8")).hexdigest() + return root / "endpoints" / namespace diff --git a/src/modelscope_hub/_download.py b/src/modelscope_hub/_download.py index 069bab1..4afc950 100644 --- a/src/modelscope_hub/_download.py +++ b/src/modelscope_hub/_download.py @@ -37,6 +37,7 @@ from tqdm.auto import tqdm from urllib3.util.retry import Retry +from ._cache_paths import endpoint_cache_root from .constants import ( DOWNLOAD_CHUNK_SIZE, DOWNLOAD_PARALLEL_THRESHOLD, @@ -1024,7 +1025,7 @@ def _repo_cache_dir_path( cache_dir: Path | None = None, ) -> Path: """Compute the repo cache directory path without creating it.""" - base = cache_dir or self._config.cache_dir + base = endpoint_cache_root(cache_dir or self._config.cache_dir, self._client.endpoint) segment = f"{repo_type}s" if not repo_type.endswith("s") else repo_type safe_id = repo_id.replace("/", "--") return base / segment / safe_id @@ -1059,7 +1060,9 @@ def _find_legacy_repo_dir( Returns the first existing, non-empty candidate, or ``None`` when the cache is clean (so the caller falls back to the new layout). """ - base = cache_dir or self._config.cache_dir + base = endpoint_cache_root(cache_dir or self._config.cache_dir, self._client.endpoint) + if base != Path(cache_dir or self._config.cache_dir): + return None segment = f"{repo_type}s" if not repo_type.endswith("s") else repo_type parts = repo_id.split("/", 1) if len(parts) != 2: @@ -1096,7 +1099,7 @@ def _lock_path( (repo type + repo id + optional file path), but store only a stable SHA-256 digest in the basename. """ - base = cache_dir or self._config.cache_dir + base = endpoint_cache_root(cache_dir or self._config.cache_dir, self._client.endpoint) scope = "file" if file_path is not None else "repo" key = "\0".join( ( diff --git a/src/modelscope_hub/api.py b/src/modelscope_hub/api.py index 2718b29..0ce1cdf 100644 --- a/src/modelscope_hub/api.py +++ b/src/modelscope_hub/api.py @@ -2454,6 +2454,7 @@ def verify_cache( revision=revision, cache_dir=cache_dir, local_dir=local_dir, + endpoint=self._config.endpoint, ) files = self.list_repo_files(repo_id, rt, revision=resolved_revision, recursive=True) expected = {file.path: file.sha256 for file in files if not file.is_dir and file.path} @@ -2464,6 +2465,7 @@ def verify_cache( revision=resolved_revision, cache_dir=Path(cache_dir) if cache_dir else None, local_dir=Path(local_dir) if local_dir else None, + endpoint=self._config.endpoint, ) def scan_cache(self, cache_dir: str | Path | None = None) -> CacheInfo: @@ -2487,7 +2489,10 @@ def scan_cache(self, cache_dir: str | Path | None = None) -> CacheInfo: >>> [r.repo_id for r in info.repos][:3] ['alice/llama-7b', 'bob/imagenet', 'carol/whisper-base'] """ - return _scan_cache(Path(cache_dir) if cache_dir else None) + return _scan_cache( + Path(cache_dir) if cache_dir else None, + endpoint=self._config.endpoint, + ) def clear_cache( self, @@ -2530,4 +2535,5 @@ def clear_cache( cache_dir=Path(cache_dir) if cache_dir else None, repo_type=rt_value, repo_id=repo_id, + endpoint=self._config.endpoint, ) diff --git a/tests/test_endpoint_cache_namespace.py b/tests/test_endpoint_cache_namespace.py new file mode 100644 index 0000000..777dc51 --- /dev/null +++ b/tests/test_endpoint_cache_namespace.py @@ -0,0 +1,100 @@ +"""Cache entries must remain bound to the registry that supplied them.""" + +from __future__ import annotations + +import hashlib +from pathlib import Path + +from modelscope_hub._download import DownloadManager +from modelscope_hub.api import HubApi +from modelscope_hub.constants import DEFAULT_ENDPOINT +from modelscope_hub.types import FileInfo + + +def _replace_download_with_endpoint_marker(monkeypatch) -> None: + def fake_download( + self: DownloadManager, + repo_id: str, + repo_type: str, + file_path: str, + revision: str, + target: Path, + **kwargs, + ) -> Path: + target.parent.mkdir(parents=True, exist_ok=True) + target.write_text(self._client.endpoint) + return target + + monkeypatch.setattr(DownloadManager, "_download_with_resume", fake_download) + + +def test_custom_endpoints_cannot_reuse_each_others_artifacts(tmp_path, monkeypatch): + _replace_download_with_endpoint_marker(monkeypatch) + first = HubApi(endpoint="https://registry-a.example") + second = HubApi(endpoint="https://registry-b.example") + + first_path = first.download_file("owner/repo", "model", "config.json", cache_dir=tmp_path) + second_path = second.download_file("owner/repo", "model", "config.json", cache_dir=tmp_path) + + assert first_path != second_path + assert first_path.read_text() == "https://registry-a.example" + assert second_path.read_text() == "https://registry-b.example" + assert first_path.relative_to(tmp_path).parts[0] == "endpoints" + assert second_path.relative_to(tmp_path).parts[0] == "endpoints" + + +def test_default_endpoint_keeps_existing_cache_layout(tmp_path, monkeypatch): + _replace_download_with_endpoint_marker(monkeypatch) + api = HubApi(endpoint=DEFAULT_ENDPOINT) + + path = api.download_file("owner/repo", "model", "config.json", cache_dir=tmp_path) + + assert path == tmp_path / "models" / "owner--repo" / "snapshots" / "master" / "config.json" + + +def test_custom_endpoint_does_not_reuse_unscoped_legacy_cache(tmp_path, monkeypatch): + _replace_download_with_endpoint_marker(monkeypatch) + legacy = tmp_path / "models" / "owner" / "repo" + legacy.mkdir(parents=True) + (legacy / "config.json").write_text("default-registry-bytes") + api = HubApi(endpoint="https://registry.example") + + path = api.download_file("owner/repo", "model", "config.json", cache_dir=tmp_path) + + assert path != legacy / "config.json" + assert path.read_text() == "https://registry.example" + + +def test_cache_management_is_scoped_to_current_endpoint(tmp_path, monkeypatch): + _replace_download_with_endpoint_marker(monkeypatch) + first = HubApi(endpoint="https://registry-a.example") + second = HubApi(endpoint="https://registry-b.example") + first_path = first.download_file("owner/repo", "model", "config.json", cache_dir=tmp_path) + second_path = second.download_file("owner/repo", "model", "config.json", cache_dir=tmp_path) + + first_info = first.scan_cache(cache_dir=tmp_path) + second_info = second.scan_cache(cache_dir=tmp_path) + first.clear_cache(cache_dir=tmp_path, repo_type="model", repo_id="owner/repo") + + assert [repo.repo_id for repo in first_info.repos] == ["owner/repo"] + assert [repo.repo_id for repo in second_info.repos] == ["owner/repo"] + assert not first_path.exists() + assert second_path.exists() + + +def test_cache_verification_uses_current_endpoint_namespace(tmp_path, monkeypatch): + _replace_download_with_endpoint_marker(monkeypatch) + api = HubApi(endpoint="https://registry.example") + path = api.download_file("owner/repo", "model", "config.json", cache_dir=tmp_path) + digest = hashlib.sha256(path.read_bytes()).hexdigest() + monkeypatch.setattr( + api, + "list_repo_files", + lambda *args, **kwargs: [FileInfo(path="config.json", sha256=digest)], + ) + + result = api.verify_cache("owner/repo", "model", revision="master", cache_dir=tmp_path) + + assert result.verified_path == str(path.parent) + assert result.checked_count == 1 + assert not result.mismatches