Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
19 changes: 14 additions & 5 deletions src/modelscope_hub/_cache_manager.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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
Expand All @@ -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))
Expand Down Expand Up @@ -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.

Expand All @@ -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.
Expand Down Expand Up @@ -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(
Expand All @@ -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
Expand Down Expand Up @@ -275,14 +280,18 @@ 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()
if not root.is_dir():
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"
Expand Down
24 changes: 24 additions & 0 deletions src/modelscope_hub/_cache_paths.py
Original file line number Diff line number Diff line change
@@ -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
9 changes: 6 additions & 3 deletions src/modelscope_hub/_download.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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:
Expand Down Expand Up @@ -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(
(
Expand Down
8 changes: 7 additions & 1 deletion src/modelscope_hub/api.py
Original file line number Diff line number Diff line change
Expand Up @@ -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}
Expand All @@ -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:
Expand All @@ -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,
Expand Down Expand Up @@ -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,
)
100 changes: 100 additions & 0 deletions tests/test_endpoint_cache_namespace.py
Original file line number Diff line number Diff line change
@@ -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