diff --git a/src/modelscope_hub/_download.py b/src/modelscope_hub/_download.py index 30e0609..069bab1 100644 --- a/src/modelscope_hub/_download.py +++ b/src/modelscope_hub/_download.py @@ -505,7 +505,7 @@ def download_file( repo_id: str, repo_type: str, file_path: str, - revision: str = "master", + revision: str | None = None, cache_dir: Path | None = None, local_dir: Path | None = None, force: bool = False, @@ -558,6 +558,17 @@ def download_file( Path Absolute path to the downloaded (or cached) file on disk. """ + if local_files_only: + return self._resolve_cached_file( + repo_id, + repo_type, + file_path, + revision=revision, + cache_dir=cache_dir, + local_dir=local_dir, + ) + + effective_revision = revision or "master" if local_dir is not None: target = Path(local_dir) / file_path else: @@ -567,17 +578,7 @@ def download_file( target = legacy / file_path else: root = self._repo_cache_dir(repo_id, repo_type, cache_dir) - target = root / "snapshots" / revision / file_path - - if local_files_only: - if target.exists(): - return target - raise CacheNotFound( - "Cannot find the requested files in the cached path and outgoing" - " traffic has been disabled. To enable look-ups and downloads" - " online, set 'local_files_only' to False.", - cache_dir=str(target.parent), - ) + target = root / "snapshots" / effective_revision / file_path if not force and self._cache_hit(target, expected_sha256): return target @@ -596,7 +597,7 @@ def download_file( repo_id, repo_type, file_path, - revision, + effective_revision, target, file_size=file_size, user_agent=user_agent, @@ -627,7 +628,7 @@ def download_repo( self, repo_id: str, repo_type: str, - revision: str = "master", + revision: str | None = None, cache_dir: Path | None = None, local_dir: Path | None = None, allow_patterns: list[str] | None = None, @@ -669,6 +670,16 @@ def download_repo( Path Absolute path to the snapshot/local directory. """ + if local_files_only: + return self._resolve_cached_snapshot( + repo_id, + repo_type, + revision=revision, + cache_dir=cache_dir, + local_dir=local_dir, + ) + + effective_revision = revision or "master" if local_dir is not None: output_dir = ensure_dir(Path(local_dir)) else: @@ -683,37 +694,26 @@ def download_repo( local_dir = legacy else: root = self._repo_cache_dir(repo_id, repo_type, cache_dir) - output_dir = ensure_dir(root / "snapshots" / revision) - - if local_files_only: - if any(output_dir.iterdir()): - logger.warning("Cannot confirm the cached file is for revision: %s", revision) - return output_dir - raise CacheNotFound( - "Cannot find the requested files in the cached path and outgoing" - " traffic has been disabled. To enable look-ups and downloads" - " online, set 'local_files_only' to False.", - cache_dir=str(output_dir), - ) + output_dir = ensure_dir(root / "snapshots" / effective_revision) if repo_type in ("skill", "skills"): return self._download_archive( repo_id=repo_id, repo_type=repo_type, - revision=revision, + revision=effective_revision, output_dir=output_dir, ) if repo_type in ("dataset", "datasets"): files = self._client.list_dataset_files_paginated( repo_id=repo_id, - revision=revision, + revision=effective_revision, ) else: files = self._client.list_repo_files( repo_id=repo_id, repo_type=repo_type, - revision=revision, + revision=effective_revision, recursive=True, ) @@ -735,10 +735,10 @@ def download_repo( download_items.append((path, sha256, size)) if not download_items: - logger.info("No files to download for %s@%s", repo_id, revision) + logger.info("No files to download for %s@%s", repo_id, effective_revision) return output_dir - logger.info("Downloading %d files from %s@%s", len(download_items), repo_id, revision) + logger.info("Downloading %d files from %s@%s", len(download_items), repo_id, effective_revision) errors: list[str] = [] with ThreadPoolExecutor(max_workers=max_workers) as executor: @@ -748,7 +748,7 @@ def download_repo( repo_id=repo_id, repo_type=repo_type, file_path=fp, - revision=revision, + revision=effective_revision, cache_dir=cache_dir, local_dir=local_dir, expected_sha256=sha256, @@ -843,17 +843,200 @@ def _cache_hit(target: Path, expected_sha256: str | None) -> bool: logger.debug("Cache hit: %s", target) return True - def _repo_cache_dir( + @staticmethod + def _dir_has_entries(path: Path) -> bool: + """Return True when a directory exists and contains at least one entry.""" + if not path.is_dir(): + return False + try: + next(path.iterdir()) + return True + except (StopIteration, OSError): + return False + + @classmethod + def _cached_snapshot_dirs(cls, snapshots_dir: Path) -> list[Path]: + """Return non-empty cached snapshot directories sorted by revision name.""" + if not snapshots_dir.is_dir(): + return [] + try: + children = sorted(snapshots_dir.iterdir(), key=lambda p: p.name) + except OSError: + return [] + return [path for path in children if cls._dir_has_entries(path)] + + @classmethod + def _describe_cached_revisions(cls, snapshots_dir: Path) -> str: + """Describe the revisions physically present in a snapshots directory.""" + if not snapshots_dir.is_dir(): + return f"No snapshots directory found at: {snapshots_dir}" + try: + children = sorted((path for path in snapshots_dir.iterdir() if path.is_dir()), key=lambda p: p.name) + except OSError as exc: + return f"Unable to inspect cached revisions under {snapshots_dir}: {exc}" + if not children: + return f"No cached revisions found under: {snapshots_dir}" + shown = [f"{path.name}{' (empty)' if not cls._dir_has_entries(path) else ''}" for path in children[:10]] + if len(children) > 10: + shown.append(f"... {len(children) - 10} more") + return f"Cached revisions under {snapshots_dir}: {', '.join(shown)}" + + def _resolve_cached_snapshot( + self, + repo_id: str, + repo_type: str, + *, + revision: str | None, + cache_dir: Path | None, + local_dir: Path | None, + ) -> Path: + """Resolve a cached snapshot without touching the network or creating directories.""" + if local_dir is not None: + output_dir = Path(local_dir) + if self._dir_has_entries(output_dir): + return output_dir + raise CacheNotFound( + "Cannot find cached snapshot in the requested local directory and outgoing traffic has been disabled. " + f"Local directory: {output_dir}", + cache_dir=str(output_dir), + ) + + legacy = self._find_legacy_repo_dir(repo_id, repo_type, cache_dir) + if legacy is not None: + logger.info("Found legacy cache at %s, reusing.", legacy) + return legacy + + root = self._repo_cache_dir_path(repo_id, repo_type, cache_dir) + snapshots = root / "snapshots" + if revision: + output_dir = snapshots / revision + if self._dir_has_entries(output_dir): + return output_dir + raise CacheNotFound( + f"Cannot find cached snapshot for {repo_type} repo '{repo_id}' at revision '{revision}' and outgoing " + f"traffic has been disabled. {self._describe_cached_revisions(snapshots)}", + cache_dir=str(output_dir), + ) + + for default_name in ("master", "main"): + default_snapshot = snapshots / default_name + if self._dir_has_entries(default_snapshot): + return default_snapshot + + candidates = self._cached_snapshot_dirs(snapshots) + if len(candidates) == 1: + logger.warning( + "Using cached revision %s because no revision was provided and outgoing traffic is disabled.", + candidates[0].name, + ) + return candidates[0] + if candidates: + revisions = ", ".join(path.name for path in candidates[:10]) + if len(candidates) > 10: + revisions += f", ... {len(candidates) - 10} more" + raise CacheNotFound( + f"Cannot determine which cached snapshot to use for {repo_type} repo '{repo_id}' because no revision " + f"was provided and outgoing traffic has been disabled. Available cached revisions: {revisions}. " + "Pass 'revision' to select one.", + cache_dir=str(snapshots), + ) + raise CacheNotFound( + f"Cannot find cached snapshot for {repo_type} repo '{repo_id}' and outgoing traffic has been disabled. " + f"{self._describe_cached_revisions(snapshots)}", + cache_dir=str(snapshots), + ) + + def _resolve_cached_file( + self, + repo_id: str, + repo_type: str, + file_path: str, + *, + revision: str | None, + cache_dir: Path | None, + local_dir: Path | None, + ) -> Path: + """Resolve a cached file without touching the network or creating directories.""" + if local_dir is not None: + target = Path(local_dir) / file_path + if target.exists(): + return target + raise CacheNotFound( + f"Cannot find cached file '{file_path}' in the requested local directory and outgoing traffic has " + f"been disabled. Local directory: {Path(local_dir)}", + cache_dir=str(target.parent), + ) + + legacy = self._find_legacy_repo_dir(repo_id, repo_type, cache_dir) + if legacy is not None: + target = legacy / file_path + if target.exists(): + return target + + root = self._repo_cache_dir_path(repo_id, repo_type, cache_dir) + snapshots = root / "snapshots" + if revision: + target = snapshots / revision / file_path + if target.exists(): + return target + raise CacheNotFound( + f"Cannot find cached file '{file_path}' for {repo_type} repo '{repo_id}' at revision '{revision}' " + f"and outgoing traffic has been disabled. {self._describe_cached_revisions(snapshots)}", + cache_dir=str(target.parent), + ) + + for default_name in ("master", "main"): + target = snapshots / default_name / file_path + if target.exists(): + return target + + matches = [] + for snapshot in self._cached_snapshot_dirs(snapshots): + candidate = snapshot / file_path + if candidate.exists(): + matches.append(candidate) + if len(matches) == 1: + logger.warning( + "Using cached file from revision %s because no revision was provided and outgoing traffic is disabled.", + matches[0].relative_to(snapshots).parts[0], + ) + return matches[0] + if matches: + revisions = ", ".join(path.relative_to(snapshots).parts[0] for path in matches[:10]) + if len(matches) > 10: + revisions += f", ... {len(matches) - 10} more" + raise CacheNotFound( + f"Cannot determine which cached file to use for {repo_type} repo '{repo_id}' because no revision was " + f"provided and outgoing traffic has been disabled. File '{file_path}' exists in cached revisions: " + f"{revisions}. Pass 'revision' to select one.", + cache_dir=str(snapshots), + ) + raise CacheNotFound( + f"Cannot find cached file '{file_path}' for {repo_type} repo '{repo_id}' and outgoing traffic has been " + f"disabled. {self._describe_cached_revisions(snapshots)}", + cache_dir=str(snapshots), + ) + + def _repo_cache_dir_path( self, repo_id: str, repo_type: str, cache_dir: Path | None = None, ) -> Path: - """Compute the cache directory for a given repo.""" + """Compute the repo cache directory path without creating it.""" base = cache_dir or self._config.cache_dir segment = f"{repo_type}s" if not repo_type.endswith("s") else repo_type safe_id = repo_id.replace("/", "--") - return ensure_dir(base / segment / safe_id) + return base / segment / safe_id + + def _repo_cache_dir( + self, + repo_id: str, + repo_type: str, + cache_dir: Path | None = None, + ) -> Path: + """Compute the cache directory for a given repo.""" + return ensure_dir(self._repo_cache_dir_path(repo_id, repo_type, cache_dir)) def _find_legacy_repo_dir( self, diff --git a/src/modelscope_hub/api.py b/src/modelscope_hub/api.py index 071d449..61cf5fe 100644 --- a/src/modelscope_hub/api.py +++ b/src/modelscope_hub/api.py @@ -1376,7 +1376,7 @@ def download_file( repo_id=repo_id, repo_type=str(rt), file_path=file_path, - revision=revision or "master", + revision=revision, cache_dir=Path(cache_dir) if cache_dir else None, local_dir=Path(local_dir) if local_dir else None, force=force, @@ -1459,7 +1459,7 @@ def download_repo( return self.downloader.download_repo( repo_id=repo_id, repo_type=str(rt), - revision=revision or "master", + revision=revision, cache_dir=Path(cache_dir) if cache_dir else None, local_dir=Path(local_dir) if local_dir else None, allow_patterns=allow_patterns, diff --git a/src/modelscope_hub/cli/login.py b/src/modelscope_hub/cli/login.py index 3983839..ced39c1 100644 --- a/src/modelscope_hub/cli/login.py +++ b/src/modelscope_hub/cli/login.py @@ -10,7 +10,7 @@ import getpass from argparse import SUPPRESS -from .base import CLICommand, SubParsers, error, info, make_api, success +from .base import CLICommand, SubParsers, error, info, make_api, success, warn from .compat import add_subcmd_token_endpoint @@ -100,3 +100,5 @@ def execute(self) -> None: info(f"id : {user.id if user.id is not None else '-'}") if user.description: info(f"description: {user.description}") + if not user.username: + warn("Username is missing from the server response; showing '-' for username.") diff --git a/src/modelscope_hub/cli/main.py b/src/modelscope_hub/cli/main.py index d05050c..88a6933 100644 --- a/src/modelscope_hub/cli/main.py +++ b/src/modelscope_hub/cli/main.py @@ -42,6 +42,7 @@ from .mcp import McpCommand from .repo import CreateCommand, DeleteCommand, InfoCommand, ListCommand, RepoCommand from .secret import SecretCommand +from .studio import StudioCommand from .upload import UploadCommand # All top-level commands in registration order. Adding a new command means @@ -61,6 +62,7 @@ LogsCommand, SettingsCommand, SecretCommand, + StudioCommand, McpCommand, CacheCommand, AgentCommand, @@ -69,6 +71,13 @@ # Plugin entry-point group name _PLUGIN_GROUP = "modelscope_hub.cli_plugins" +# ``studio`` used to be contributed by the umbrella SDK. It is now built into +# this package so hub-only installs can manage Studio spaces, but older SDK +# wheels still advertise the same plugin. Silently skipping that known legacy +# plugin avoids noisy warnings on every command while preserving warnings for +# genuinely unexpected collisions. +_KNOWN_BUILTIN_PLUGIN_MIGRATIONS = frozenset({"studio"}) + # --------------------------------------------------------------------------- # Invocation identity @@ -299,6 +308,13 @@ def _discover_plugins(subparsers: Action) -> None: cmd_cls = ep.load() name = getattr(cmd_cls, "name", ep.name) if registered is not None and name in registered: + if name in _KNOWN_BUILTIN_PLUGIN_MIGRATIONS: + log.debug( + "Skipping legacy CLI plugin %r: command %r is built in.", + ep.name, + name, + ) + continue log.warning( "Skipping CLI plugin %r: command %r is already registered.", ep.name, diff --git a/src/modelscope_hub/cli/studio.py b/src/modelscope_hub/cli/studio.py new file mode 100644 index 0000000..ba825bd --- /dev/null +++ b/src/modelscope_hub/cli/studio.py @@ -0,0 +1,258 @@ +"""``ms studio`` command group — manage Studio runtime resources. + +Historically this command group lived in the umbrella ``modelscope`` SDK and was +registered into this CLI as a plugin. The console scripts are now owned by +``modelscope-hub`` itself, so the Studio management surface must be available +from the hub package too (including hub-only installs such as ``ms-hub``). +""" + +from __future__ import annotations + +import json +from argparse import SUPPRESS +from typing import Any + +from ..constants import RepoType +from .base import CLICommand, SubParsers, info, make_api, parse_kv_pairs, success + +_LOG_TYPES = ("run", "build") +_STUDIO_SDK_TYPES = ("gradio", "streamlit", "docker", "static") +_SETTINGS_FIELDS = ( + "display_name", + "description", + "license", + "cover_image", + "sdk_type", + "sdk_version", + "base_image", + "hardware", + "private", +) + + +class StudioCommand(CLICommand): + """Top-level dispatcher for ``studio`` subcommands.""" + + @staticmethod + def register(subparsers: SubParsers) -> None: + parser = subparsers.add_parser("studio", help="Manage ModelScope Studio spaces.") + _add_visible_auth_args(parser) + actions = parser.add_subparsers(dest="studio_action", metavar="ACTION") + actions.required = True + + _StudioDeploy.register(actions) + _StudioStop.register(actions) + _StudioLogs.register(actions) + _StudioSettings.register(actions) + _StudioSecret.register(actions) + + parser.set_defaults(_command=StudioCommand) + + def execute(self) -> None: + leaf = getattr(self.args, "_studio_leaf", None) + if leaf is None: # pragma: no cover - argparse enforces this + raise SystemExit("No studio subcommand specified. Run 'modelscope studio --help'.") + leaf(self.args).execute() + + +# Backward-compatible name used by the umbrella SDK's historical module. +StudioCMD = StudioCommand + + +def _add_visible_auth_args(parser) -> None: + parser.add_argument("--token", dest="subcmd_token", default=None, help="Optional access token.") + parser.add_argument("--endpoint", dest="subcmd_endpoint", default=None, help="ModelScope server endpoint.") + + +def _add_studio_id(parser) -> None: + parser.add_argument("studio_id", help="Studio ID in the form 'owner/name'.") + _add_leaf_auth_args(parser) + + +def _add_leaf_auth_args(parser) -> None: + """Accept auth flags after a studio action without overwriting group flags.""" + parser.add_argument("--token", dest="subcmd_token", default=SUPPRESS, help=SUPPRESS) + parser.add_argument("--endpoint", dest="subcmd_endpoint", default=SUPPRESS, help=SUPPRESS) + + +class _StudioDeploy(CLICommand): + @staticmethod + def register(subparsers: SubParsers) -> None: + p = subparsers.add_parser("deploy", help="Deploy (re-pull and rebuild) a Studio space.") + _add_studio_id(p) + p.set_defaults(_command=StudioCommand, _studio_leaf=_StudioDeploy) + + def execute(self) -> None: + api = make_api(self.args) + data = api.deploy_repo(self.args.studio_id, RepoType.STUDIO) + success(f"Deploy triggered for studio {self.args.studio_id}.") + _print_status(data) + + +class _StudioStop(CLICommand): + @staticmethod + def register(subparsers: SubParsers) -> None: + p = subparsers.add_parser("stop", help="Stop a running Studio space.") + _add_studio_id(p) + p.set_defaults(_command=StudioCommand, _studio_leaf=_StudioStop) + + def execute(self) -> None: + api = make_api(self.args) + data = api.stop_repo(self.args.studio_id, RepoType.STUDIO) + success(f"Stop triggered for studio {self.args.studio_id}.") + _print_status(data) + + +class _StudioLogs(CLICommand): + @staticmethod + def register(subparsers: SubParsers) -> None: + p = subparsers.add_parser("logs", help="Fetch Studio runtime or build logs.") + _add_studio_id(p) + p.add_argument("--type", "--log-type", dest="log_type", choices=_LOG_TYPES, default="run") + p.add_argument("--keyword", default=None, help="Optional keyword to filter log lines.") + p.add_argument("--page", "--page-num", dest="page_num", type=int, default=1) + p.add_argument("--page-size", dest="page_size", type=int, default=100) + p.add_argument("--start-timestamp", dest="start_timestamp", type=int, default=None) + p.add_argument("--end-timestamp", dest="end_timestamp", type=int, default=None) + p.set_defaults(_command=StudioCommand, _studio_leaf=_StudioLogs) + + def execute(self) -> None: + api = make_api(self.args) + payload = api.get_repo_logs( + self.args.studio_id, + RepoType.STUDIO, + log_type=self.args.log_type, + page_num=self.args.page_num, + page_size=self.args.page_size, + keyword=self.args.keyword, + start_timestamp=self.args.start_timestamp, + end_timestamp=self.args.end_timestamp, + ) + _print_logs(payload, page_num=self.args.page_num, page_size=self.args.page_size) + + +class _StudioSettings(CLICommand): + @staticmethod + def register(subparsers: SubParsers) -> None: + p = subparsers.add_parser("settings", help="Update Studio settings.") + _add_studio_id(p) + p.add_argument("settings", nargs="*", help="Optional key=value settings.") + p.add_argument("--display-name", dest="display_name", default=None, help="Studio display name.") + p.add_argument("--description", default=None, help="Studio description.") + p.add_argument("--license", default=None, help="Studio license.") + p.add_argument("--cover-image", dest="cover_image", default=None, help="Studio cover image URL.") + p.add_argument("--sdk-type", dest="sdk_type", choices=_STUDIO_SDK_TYPES, default=None) + p.add_argument("--sdk-version", dest="sdk_version", default=None) + p.add_argument("--base-image", dest="base_image", default=None) + p.add_argument("--hardware", default=None) + visibility = p.add_mutually_exclusive_group() + visibility.add_argument("--private", dest="private", action="store_const", const=True, default=None) + visibility.add_argument("--public", dest="private", action="store_const", const=False) + p.set_defaults(_command=StudioCommand, _studio_leaf=_StudioSettings) + + def execute(self) -> None: + settings = _collect_settings(self.args) + if not settings: + raise ValueError( + "No setting specified. Provide key=value or one of: --display-name, --description, --license, " + "--cover-image, --sdk-type, --sdk-version, --base-image, --hardware, --private/--public." + ) + api = make_api(self.args) + data = api.update_repo_settings(self.args.studio_id, RepoType.STUDIO, **settings) + success(f"Updated settings for studio {self.args.studio_id}: {', '.join(sorted(settings))}.") + if data: + info(json.dumps(data, ensure_ascii=False, indent=2, default=str)) + + +class _StudioSecret(CLICommand): + @staticmethod + def register(subparsers: SubParsers) -> None: + parser = subparsers.add_parser("secret", help="Manage Studio environment variables (secrets).") + actions = parser.add_subparsers(dest="secret_action", metavar="ACTION") + actions.required = True + + list_p = actions.add_parser("list", help="List secret keys.") + _add_studio_id(list_p) + list_p.set_defaults(_command=StudioCommand, _studio_leaf=_StudioSecret) + + add_p = actions.add_parser("add", help="Add a secret.") + _add_studio_id(add_p) + add_p.add_argument("key") + add_p.add_argument("value") + add_p.set_defaults(_command=StudioCommand, _studio_leaf=_StudioSecret) + + update_p = actions.add_parser("update", help="Update an existing secret.") + _add_studio_id(update_p) + update_p.add_argument("key") + update_p.add_argument("value") + update_p.set_defaults(_command=StudioCommand, _studio_leaf=_StudioSecret) + + delete_p = actions.add_parser("delete", help="Delete a secret.") + _add_studio_id(delete_p) + delete_p.add_argument("key") + delete_p.set_defaults(_command=StudioCommand, _studio_leaf=_StudioSecret) + + def execute(self) -> None: + api = make_api(self.args) + action = self.args.secret_action + if action == "list": + secrets = api.list_secrets(self.args.studio_id, RepoType.STUDIO) + if not secrets: + info("(no secrets)") + return + for item in secrets: + info(str(item.get("key") if isinstance(item, dict) else item)) + return + if action == "add": + api.add_secret(self.args.studio_id, self.args.key, self.args.value, RepoType.STUDIO) + success(f"Secret {self.args.key!r} added.") + return + if action == "update": + api.update_secret(self.args.studio_id, self.args.key, self.args.value, RepoType.STUDIO) + success(f"Secret {self.args.key!r} updated.") + return + if action == "delete": + api.delete_secret(self.args.studio_id, self.args.key, RepoType.STUDIO) + success(f"Secret {self.args.key!r} deleted.") + return + raise ValueError(f"Unknown secret subcommand: {action}") + + +def _collect_settings(args) -> dict[str, Any]: + settings: dict[str, Any] = parse_kv_pairs(getattr(args, "settings", []) or []) + for field in _SETTINGS_FIELDS: + value = getattr(args, field, None) + if value is not None: + settings[field] = value + return settings + + +def _print_status(data: object) -> None: + if not data: + return + if isinstance(data, dict): + status = data.get("status") or data.get("Status") + if status: + info(f"Status: {status}") + return + info(json.dumps(data, ensure_ascii=False, indent=2, default=str)) + + +def _print_logs(payload: object, *, page_num: int, page_size: int) -> None: + if not isinstance(payload, dict): + info(str(payload)) + return + logs = payload.get("logs") + if logs is None: + info(json.dumps(payload, ensure_ascii=False, indent=2, default=str)) + return + for entry in logs: + if isinstance(entry, dict): + ts = entry.get("timestamp") or entry.get("time") or "" + msg = entry.get("content") or entry.get("message") or "" + info(f"[{ts}] {msg}" if ts else str(msg)) + else: + info(str(entry)) + total = payload.get("total") + if total is not None: + info(f"-- page {page_num} (size {page_size}), total {total} --") diff --git a/src/modelscope_hub/compat/file_download.py b/src/modelscope_hub/compat/file_download.py index 58b1371..5d81739 100644 --- a/src/modelscope_hub/compat/file_download.py +++ b/src/modelscope_hub/compat/file_download.py @@ -13,7 +13,6 @@ from ..api import HubApi from ..constants import RepoType from ..errors import AuthenticationError, NotExistError, PermissionDeniedError -from .constants import DEFAULT_DATASET_REVISION def _resolve_legacy_paths( @@ -123,7 +122,7 @@ def dataset_file_download( dataset_id, repo_type=RepoType.DATASET, file_path=file_path, - revision=revision or DEFAULT_DATASET_REVISION, + revision=revision, cache_dir=effective_cache, local_dir=effective_local, local_files_only=local_files_only, diff --git a/src/modelscope_hub/compat/snapshot_download.py b/src/modelscope_hub/compat/snapshot_download.py index 7990508..3981972 100644 --- a/src/modelscope_hub/compat/snapshot_download.py +++ b/src/modelscope_hub/compat/snapshot_download.py @@ -17,7 +17,6 @@ from ..constants import RepoType from ..errors import AuthenticationError, NotExistError, PermissionDeniedError from ..utils.patterns import normalize_patterns -from .constants import DEFAULT_DATASET_REVISION from .file_download import _resolve_legacy_paths if TYPE_CHECKING: @@ -150,7 +149,7 @@ def dataset_snapshot_download( result = api.download_repo( dataset_id, repo_type=RepoType.DATASET, - revision=revision or DEFAULT_DATASET_REVISION, + revision=revision, cache_dir=effective_cache, local_dir=effective_local, allow_patterns=include, diff --git a/src/modelscope_hub/config.py b/src/modelscope_hub/config.py index 9fbbadc..87d7d31 100644 --- a/src/modelscope_hub/config.py +++ b/src/modelscope_hub/config.py @@ -14,6 +14,8 @@ from __future__ import annotations import os +import stat +import uuid import warnings from dataclasses import dataclass, field from pathlib import Path @@ -46,6 +48,57 @@ USER_INFO_FILE_NAME, ) +# Credentials are private even when they are not bearer tokens: ``user`` stores +# username/email and ``session`` is a stable install identifier. Do not rely on +# process umask for files under ``~/.modelscope/credentials``. +_PRIVATE_DIR_MODE = stat.S_IRWXU +_PRIVATE_FILE_MODE = stat.S_IRUSR | stat.S_IWUSR + + +def _chmod_private(path: Path, mode: int) -> None: + """Best-effort chmod used for credential files and directories.""" + try: + path.chmod(mode) + except OSError: + # Some filesystems/OSes do not implement POSIX modes fully. Credential + # reads should still work; tests assert exact modes only on POSIX. + pass + + +def _write_private_bytes(path: Path, data: bytes) -> None: + """Atomically write *data* with a private initial mode. + + Creating the temporary file with ``0o600`` avoids the short window where a + normal ``open``/``write_text`` would respect a permissive umask (commonly + yielding ``0644``) before a later chmod could tighten permissions. + """ + tmp = path.with_name(f".{path.name}.tmp-{os.getpid()}-{uuid.uuid4().hex}") + fd: int | None = None + try: + fd = os.open(tmp, os.O_WRONLY | os.O_CREAT | os.O_EXCL, _PRIVATE_FILE_MODE) + with os.fdopen(fd, "wb") as f: + fd = None + f.write(data) + _chmod_private(tmp, _PRIVATE_FILE_MODE) + os.replace(tmp, path) + _chmod_private(path, _PRIVATE_FILE_MODE) + except Exception: + if fd is not None: + try: + os.close(fd) + except OSError: + pass + try: + tmp.unlink(missing_ok=True) + except OSError: + pass + raise + + +def _write_private_text(path: Path, text: str) -> None: + """Atomically write UTF-8 text with ``0600`` permissions.""" + _write_private_bytes(path, text.encode("utf-8")) + def _expand(path: str | os.PathLike[str]) -> Path: return Path(path).expanduser().resolve() @@ -94,6 +147,7 @@ def __post_init__(self) -> None: else: self.endpoint = DEFAULT_ENDPOINT self.endpoint = self.normalize_endpoint(self.endpoint) + self._repair_credential_permissions() # Token precedence: explicit arg > MODELSCOPE_API_TOKEN env var > # persisted credential. An explicitly provided value wins even when # empty ("" means "use no token"), so an explicit override never @@ -138,9 +192,29 @@ def ensure_dirs(self) -> None: self.config_dir.mkdir(parents=True, exist_ok=True) self.cache_dir.mkdir(parents=True, exist_ok=True) self.credentials_dir.mkdir(parents=True, exist_ok=True) + _chmod_private(self.credentials_dir, _PRIVATE_DIR_MODE) except OSError as exc: # pragma: no cover - filesystem dependent raise CacheError(f"Failed to create SDK directories: {exc}") from exc + def _repair_credential_permissions(self) -> None: + """Tighten permissions for credentials written by older versions. + + This is intentionally non-creating: constructing ``HubConfig`` should + not materialise ``~/.modelscope`` for users who only use environment + variables. If a credentials directory already exists, repair its mode and + every known file inside it. + """ + try: + if not self.credentials_dir.is_dir(): + return + _chmod_private(self.credentials_dir, _PRIVATE_DIR_MODE) + for name in (*_CREDENTIAL_FILE_NAMES, SESSION_FILE_NAME): + path = self.credentials_dir / name + if path.is_file(): + _chmod_private(path, _PRIVATE_FILE_MODE) + except OSError: + pass + # ------------------------------------------------------------------ # Token persistence # ------------------------------------------------------------------ @@ -230,13 +304,10 @@ def clear_token(self) -> None: def save_cookies(self, cookies: object) -> None: """Pickle cookies to ``~/.modelscope/credentials/cookies``.""" import pickle - import stat self.ensure_dirs() path = self.credentials_dir / COOKIES_FILE_NAME - with open(path, "wb") as f: - pickle.dump(cookies, f) - path.chmod(stat.S_IRUSR | stat.S_IWUSR) + _write_private_bytes(path, pickle.dumps(cookies)) def load_cookies(self) -> Any: """Load saved cookies, returning None if absent or expired.""" @@ -261,16 +332,13 @@ def save_user_info(self, username: str, email: str) -> None: """Save ``username:email`` to ``~/.modelscope/credentials/user``.""" self.ensure_dirs() path = self.credentials_dir / USER_INFO_FILE_NAME - path.write_text(f"{username}:{email}", encoding="utf-8") + _write_private_text(path, f"{username}:{email}") def save_git_token(self, git_token: str) -> None: """Save git token to ``~/.modelscope/credentials/git_token``.""" - import stat - self.ensure_dirs() path = self.credentials_dir / GIT_TOKEN_FILE_NAME - path.write_text(git_token, encoding="utf-8") - path.chmod(stat.S_IRUSR | stat.S_IWUSR) + _write_private_text(path, git_token) def load_git_token(self) -> str | None: """Read git token from ``~/.modelscope/credentials/git_token``.""" @@ -288,19 +356,18 @@ def get_session_id(self) -> str: The session ID is persisted to ``~/.modelscope/credentials/session`` and included in the User-Agent header for telemetry. """ - import uuid as _uuid - path = self.credentials_dir / SESSION_FILE_NAME if path.is_file(): try: sid = path.read_text(encoding="utf-8").strip() if len(sid) == 32: + _chmod_private(path, _PRIVATE_FILE_MODE) return sid except OSError: pass - sid = _uuid.uuid4().hex + sid = uuid.uuid4().hex self.ensure_dirs() - path.write_text(sid, encoding="utf-8") + _write_private_text(path, sid) return sid diff --git a/src/modelscope_hub/types.py b/src/modelscope_hub/types.py index 15d7de2..4efcbeb 100644 --- a/src/modelscope_hub/types.py +++ b/src/modelscope_hub/types.py @@ -61,6 +61,55 @@ class UserInfo(_FromDictMixin): "name": "username", "avatar": "avatar_url", } + _id_keys: ClassVar[tuple[str, ...]] = ( + "id", + "Id", + "ID", + "user_id", + "UserId", + "userId", + "uid", + "Uid", + "UID", + "sub", + "Sub", + ) + _username_keys: ClassVar[tuple[str, ...]] = ( + "Username", + "username", + # Observed OIDC-style shape on newer /users/me responses: the login + # handle may be in ``name`` while ``preferred_username`` can be empty or + # a display value, so keep the same priority as get_current_username(). + "name", + "Name", + "preferred_username", + "PreferredUsername", + "preferredUsername", + "user_name", + "UserName", + "login", + "Login", + "nickname", + "Nickname", + ) + _email_keys: ClassVar[tuple[str, ...]] = ("email", "Email", "mail", "Mail") + _avatar_keys: ClassVar[tuple[str, ...]] = ( + "avatar_url", + "avatarUrl", + "AvatarUrl", + "avatar", + "Avatar", + "picture", + "Picture", + ) + _description_keys: ClassVar[tuple[str, ...]] = ( + "description", + "Description", + "bio", + "Bio", + "introduction", + "Introduction", + ) id: str | int | None = None username: str | None = None @@ -68,6 +117,35 @@ class UserInfo(_FromDictMixin): avatar_url: str | None = None description: str | None = None + @classmethod + def from_dict(cls, data: Mapping[str, Any] | None) -> UserInfo: + """Build user info from legacy ModelScope and newer OIDC-style keys.""" + if not isinstance(data, Mapping) or not data: + return cls() + return cls( + id=cls._first_non_empty(data, cls._id_keys), + username=cls._as_str_or_none(cls._first_non_empty(data, cls._username_keys)), + email=cls._as_str_or_none(cls._first_non_empty(data, cls._email_keys)), + avatar_url=cls._as_str_or_none(cls._first_non_empty(data, cls._avatar_keys)), + description=cls._as_str_or_none(cls._first_non_empty(data, cls._description_keys)), + ) + + @staticmethod + def _first_non_empty(data: Mapping[str, Any], keys: tuple[str, ...]) -> Any | None: + for key in keys: + if key not in data: + continue + value = data[key] + if value is not None and value != "": + return value + return None + + @staticmethod + def _as_str_or_none(value: Any | None) -> str | None: + if value is None: + return None + return str(value) + # --------------------------------------------------------------------------- # Repository diff --git a/tests/cli/test_entry_points.py b/tests/cli/test_entry_points.py index d2f500b..b398087 100644 --- a/tests/cli/test_entry_points.py +++ b/tests/cli/test_entry_points.py @@ -284,6 +284,23 @@ def register(subparsers): args = parser.parse_args(["download", "owner/name"]) assert args.repo_id == "owner/name" + def test_legacy_studio_plugin_collision_is_silent(self, only_plugins, plugin_warnings): + """Old ``modelscope`` wheels still advertise ``studio``; hub now owns it.""" + + class _LegacyStudio: + name = "studio" + + @staticmethod + def register(subparsers): + raise AssertionError("legacy studio plugin must be skipped") + + only_plugins(_FakeEntryPoint("studio", _LegacyStudio)) + + parser = _build_parser() + + assert plugin_warnings == [] + assert parser.parse_args(["studio", "deploy", "owner/demo"]).studio_action == "deploy" + def test_unimportable_plugin_does_not_break_the_cli(self, only_plugins, plugin_warnings): """Optional extras are legitimately absent, so this must stay quiet.""" only_plugins(_FakeEntryPoint("broken", None, boom=True)) diff --git a/tests/cli/test_login.py b/tests/cli/test_login.py index 920765d..e4a204f 100644 --- a/tests/cli/test_login.py +++ b/tests/cli/test_login.py @@ -203,8 +203,9 @@ def test_whoami_missing_fields(self, parser, mock_api, capsys): args = parser.parse_args(["whoami"]) with patch("modelscope_hub.cli.login.make_api", return_value=mock_api): WhoamiCommand(args).execute() - out = capsys.readouterr().out - assert "-" in out + captured = capsys.readouterr() + assert "-" in captured.out + assert "Username is missing" in captured.err # =================================================================== diff --git a/tests/cli/test_studio.py b/tests/cli/test_studio.py new file mode 100644 index 0000000..23b47dd --- /dev/null +++ b/tests/cli/test_studio.py @@ -0,0 +1,225 @@ +"""Tests for the built-in ``studio`` command group.""" + +from __future__ import annotations + +from unittest.mock import patch + +import pytest + +from modelscope_hub.cli.studio import StudioCommand +from modelscope_hub.constants import RepoType + +from .conftest import run_cli + + +# =================================================================== +# Parser tests +# =================================================================== +class TestStudioParser: + def test_studio_group_is_registered(self, parser): + with pytest.raises(SystemExit) as exc_info: + parser.parse_args(["studio", "--help"]) + assert exc_info.value.code == 0 + + def test_deploy(self, parser): + args = parser.parse_args(["studio", "deploy", "org/demo"]) + assert args._command is StudioCommand + assert args.studio_action == "deploy" + assert args.studio_id == "org/demo" + + def test_stop(self, parser): + args = parser.parse_args(["studio", "stop", "org/demo"]) + assert args.studio_action == "stop" + assert args.studio_id == "org/demo" + + def test_logs_options(self, parser): + args = parser.parse_args( + [ + "studio", + "logs", + "org/demo", + "--type", + "build", + "--keyword", + "ERROR", + "--page-num", + "3", + "--page-size", + "50", + "--start-timestamp", + "10", + "--end-timestamp", + "20", + ] + ) + assert args.log_type == "build" + assert args.keyword == "ERROR" + assert args.page_num == 3 + assert args.page_size == 50 + assert args.start_timestamp == 10 + assert args.end_timestamp == 20 + + def test_logs_log_type_alias(self, parser): + args = parser.parse_args(["studio", "logs", "org/demo", "--log-type", "run"]) + assert args.log_type == "run" + + def test_settings_flags(self, parser): + args = parser.parse_args( + [ + "studio", + "settings", + "org/demo", + "--display-name", + "Demo", + "--sdk-type", + "gradio", + "--private", + ] + ) + assert args.display_name == "Demo" + assert args.sdk_type == "gradio" + assert args.private is True + + def test_settings_key_value_tokens(self, parser): + args = parser.parse_args(["studio", "settings", "org/demo", "hardware=cpu", "private=true"]) + assert args.settings == ["hardware=cpu", "private=true"] + + def test_group_level_auth(self, parser): + args = parser.parse_args(["studio", "--token", "tk", "--endpoint", "https://x.cn", "deploy", "org/demo"]) + assert args.subcmd_token == "tk" + assert args.subcmd_endpoint == "https://x.cn" + + def test_leaf_level_auth(self, parser): + args = parser.parse_args(["studio", "deploy", "org/demo", "--token", "tk"]) + assert args.subcmd_token == "tk" + + def test_secret_add(self, parser): + args = parser.parse_args(["studio", "secret", "add", "org/demo", "API_KEY", "value"]) + assert args.studio_action == "secret" + assert args.secret_action == "add" + assert args.studio_id == "org/demo" + assert args.key == "API_KEY" + assert args.value == "value" + + +# =================================================================== +# Execution tests +# =================================================================== +@pytest.mark.mock_only +class TestStudioExecute: + def test_deploy(self, mock_api): + with patch("modelscope_hub.cli.studio.make_api", return_value=mock_api): + code, out, err = run_cli(["studio", "deploy", "org/demo"]) + assert code == 0 + assert "Deploy triggered" in out + mock_api.deploy_repo.assert_called_once_with("org/demo", RepoType.STUDIO) + + def test_stop(self, mock_api): + with patch("modelscope_hub.cli.studio.make_api", return_value=mock_api): + code, out, err = run_cli(["studio", "stop", "org/demo"]) + assert code == 0 + assert "Stop triggered" in out + mock_api.stop_repo.assert_called_once_with("org/demo", RepoType.STUDIO) + + def test_logs(self, mock_api, capsys): + mock_api.get_repo_logs.return_value = { + "logs": [ + {"timestamp": "t1", "content": "line1"}, + {"time": "t2", "message": "line2"}, + "line3", + ], + "total": 3, + } + with patch("modelscope_hub.cli.studio.make_api", return_value=mock_api): + code, out, err = run_cli( + [ + "studio", + "logs", + "org/demo", + "--type", + "build", + "--keyword", + "ERR", + "--page-num", + "2", + "--page-size", + "5", + "--start-timestamp", + "10", + "--end-timestamp", + "20", + ] + ) + assert code == 0 + assert "line1" in out + assert "line2" in out + assert "line3" in out + mock_api.get_repo_logs.assert_called_once_with( + "org/demo", + RepoType.STUDIO, + log_type="build", + page_num=2, + page_size=5, + keyword="ERR", + start_timestamp=10, + end_timestamp=20, + ) + + def test_settings_flags_and_key_value_tokens(self, mock_api): + with patch("modelscope_hub.cli.studio.make_api", return_value=mock_api): + code, out, err = run_cli( + [ + "studio", + "settings", + "org/demo", + "hardware=cpu", + "--display-name", + "Demo", + "--public", + ] + ) + assert code == 0 + assert "Updated settings" in out + mock_api.update_repo_settings.assert_called_once_with( + "org/demo", + RepoType.STUDIO, + hardware="cpu", + display_name="Demo", + private=False, + ) + + def test_settings_requires_at_least_one_setting(self, mock_api): + with patch("modelscope_hub.cli.studio.make_api", return_value=mock_api): + code, out, err = run_cli(["studio", "settings", "org/demo"]) + assert code == 2 + assert "No setting specified" in err + mock_api.update_repo_settings.assert_not_called() + + def test_secret_list(self, mock_api): + mock_api.list_secrets.return_value = [{"key": "API_KEY"}] + with patch("modelscope_hub.cli.studio.make_api", return_value=mock_api): + code, out, err = run_cli(["studio", "secret", "list", "org/demo"]) + assert code == 0 + assert "API_KEY" in out + mock_api.list_secrets.assert_called_once_with("org/demo", RepoType.STUDIO) + + def test_secret_add(self, mock_api): + with patch("modelscope_hub.cli.studio.make_api", return_value=mock_api): + code, out, err = run_cli(["studio", "secret", "add", "org/demo", "API_KEY", "value"]) + assert code == 0 + assert "added" in out.lower() + mock_api.add_secret.assert_called_once_with("org/demo", "API_KEY", "value", RepoType.STUDIO) + + def test_secret_update(self, mock_api): + with patch("modelscope_hub.cli.studio.make_api", return_value=mock_api): + code, out, err = run_cli(["studio", "secret", "update", "org/demo", "API_KEY", "new"]) + assert code == 0 + assert "updated" in out.lower() + mock_api.update_secret.assert_called_once_with("org/demo", "API_KEY", "new", RepoType.STUDIO) + + def test_secret_delete(self, mock_api): + with patch("modelscope_hub.cli.studio.make_api", return_value=mock_api): + code, out, err = run_cli(["studio", "secret", "delete", "org/demo", "API_KEY"]) + assert code == 0 + assert "deleted" in out.lower() + mock_api.delete_secret.assert_called_once_with("org/demo", "API_KEY", RepoType.STUDIO) diff --git a/tests/test_compat_snapshot_download.py b/tests/test_compat_snapshot_download.py index 8f5b85b..44686de 100644 --- a/tests/test_compat_snapshot_download.py +++ b/tests/test_compat_snapshot_download.py @@ -13,8 +13,12 @@ from modelscope_hub import ProgressCallback from modelscope_hub._download import DownloadManager +from modelscope_hub.compat.file_download import model_file_download from modelscope_hub.compat.snapshot_download import snapshot_download +_snapshot_download_mod = __import__("modelscope_hub.compat.snapshot_download", fromlist=["HubApi"]) +_file_download_mod = __import__("modelscope_hub.compat.file_download", fromlist=["HubApi"]) + class _DummyCallback(ProgressCallback): pass @@ -39,3 +43,34 @@ def test_progress_callbacks_default_none(self): _, kwargs = m.call_args assert kwargs["progress_callbacks"] is None + + def test_snapshot_local_files_only_does_not_probe_endpoint(self): + with ( + mock.patch.object( + DownloadManager, + "download_repo", + return_value="/tmp/snapshot", + ), + mock.patch.object( + _snapshot_download_mod.HubApi, + "resolve_endpoint_for_read", + side_effect=AssertionError("network should not be used"), + ), + ): + assert snapshot_download("owner/repo", local_files_only=True) == "/tmp/snapshot" + + def test_file_local_files_only_does_not_probe_endpoint(self): + with ( + mock.patch.object( + DownloadManager, + "download_file", + return_value="/tmp/snapshot/config.json", + ), + mock.patch.object( + _file_download_mod.HubApi, + "resolve_endpoint_for_read", + side_effect=AssertionError("network should not be used"), + ), + ): + result = model_file_download("owner/repo", "config.json", local_files_only=True) + assert result == "/tmp/snapshot/config.json" diff --git a/tests/test_credential_lifecycle.py b/tests/test_credential_lifecycle.py index 9c87cfe..4293264 100644 --- a/tests/test_credential_lifecycle.py +++ b/tests/test_credential_lifecycle.py @@ -12,6 +12,9 @@ from __future__ import annotations +import os +import stat + import pytest from modelscope_hub.api import HubApi @@ -25,6 +28,18 @@ TOKEN = "ms-token-under-test" GIT_TOKEN = "git-token-value" +_PRIVATE_FILE_MODE = 0o600 +_PRIVATE_DIR_MODE = 0o700 + + +def _mode(path) -> int: + return stat.S_IMODE(path.stat().st_mode) + + +_requires_posix_permissions = pytest.mark.skipif( + os.name == "nt", + reason="POSIX credential permission bits are not reliable on Windows.", +) @pytest.fixture(autouse=True) @@ -48,6 +63,64 @@ def fully_logged_in(home) -> HubConfig: return config +@_requires_posix_permissions +def test_credential_directory_is_private(isolated_home): + config = HubConfig(config_dir=isolated_home) + + config.ensure_dirs() + + assert _mode(config.credentials_dir) == _PRIVATE_DIR_MODE + + +@_requires_posix_permissions +def test_saved_credential_files_are_private(isolated_home): + config = fully_logged_in(isolated_home) + + expected_private = (COOKIES_FILE_NAME, GIT_TOKEN_FILE_NAME, USER_INFO_FILE_NAME, SESSION_FILE_NAME) + for name in expected_private: + assert _mode(config.credentials_dir / name) == _PRIVATE_FILE_MODE, name + + +@_requires_posix_permissions +def test_existing_session_permission_is_repaired(isolated_home): + config = HubConfig(config_dir=isolated_home) + config.ensure_dirs() + session_path = config.credentials_dir / SESSION_FILE_NAME + existing_session = "a" * 32 + session_path.write_text(existing_session, encoding="utf-8") + session_path.chmod(0o644) + + assert config.get_session_id() == existing_session + + assert _mode(session_path) == _PRIVATE_FILE_MODE + + +@_requires_posix_permissions +def test_existing_credential_permissions_are_repaired_on_config_load(isolated_home): + credentials = isolated_home / "credentials" + credentials.mkdir(parents=True) + credentials.chmod(0o755) + for name in (COOKIES_FILE_NAME, GIT_TOKEN_FILE_NAME, USER_INFO_FILE_NAME, SESSION_FILE_NAME): + path = credentials / name + path.write_text("placeholder", encoding="utf-8") + path.chmod(0o644) + + HubConfig(config_dir=isolated_home, token="") + + assert _mode(credentials) == _PRIVATE_DIR_MODE + for name in (COOKIES_FILE_NAME, GIT_TOKEN_FILE_NAME, USER_INFO_FILE_NAME, SESSION_FILE_NAME): + assert _mode(credentials / name) == _PRIVATE_FILE_MODE, name + + +def test_config_load_does_not_create_credentials_directory(isolated_home): + credentials = isolated_home / "credentials" + assert not credentials.exists() + + HubConfig(config_dir=isolated_home, token="") + + assert not credentials.exists() + + # --------------------------------------------------------------------------- # All-or-nothing teardown # --------------------------------------------------------------------------- diff --git a/tests/test_legacy_cache_detection.py b/tests/test_legacy_cache_detection.py index e76d46c..10c3e9f 100644 --- a/tests/test_legacy_cache_detection.py +++ b/tests/test_legacy_cache_detection.py @@ -7,7 +7,12 @@ from __future__ import annotations +from unittest import mock + +import pytest + from modelscope_hub.api import HubApi +from modelscope_hub.errors import CacheNotFound def _make_download_manager(): @@ -55,3 +60,59 @@ def test_empty_legacy_dir_returns_none(self, tmp_path): dm = _make_download_manager() assert dm._find_legacy_repo_dir("Qwen/Qwen3.5-4B", "model", tmp_path) is None + + +class TestLocalFilesOnlyCacheResolution: + def test_snapshot_without_revision_uses_single_cached_non_default_revision(self, tmp_path): + snapshot = tmp_path / "models" / "owner--repo" / "snapshots" / "v2.0.4" + snapshot.mkdir(parents=True) + (snapshot / "config.json").write_text("{}") + + api = HubApi() + with mock.patch.object(api.legacy, "list_repo_files", side_effect=AssertionError("network should not be used")): + result = api.download_repo("owner/repo", "model", cache_dir=tmp_path, local_files_only=True) + + assert result == snapshot + assert not (tmp_path / "models" / "owner--repo" / "snapshots" / "master").exists() + + def test_file_without_revision_uses_single_cached_non_default_revision(self, tmp_path): + cached_file = tmp_path / "models" / "owner--repo" / "snapshots" / "v2.0.4" / "weights" / "model.bin" + cached_file.parent.mkdir(parents=True) + cached_file.write_bytes(b"cached") + + api = HubApi() + with mock.patch.object( + api.downloader, + "_download_with_resume", + side_effect=AssertionError("network should not be used"), + ): + result = api.download_file( + "owner/repo", + "model", + "weights/model.bin", + cache_dir=tmp_path, + local_files_only=True, + ) + + assert result == cached_file + assert not (tmp_path / "models" / "owner--repo" / "snapshots" / "master").exists() + + def test_cache_not_found_mentions_available_cached_revisions(self, tmp_path): + snapshot = tmp_path / "models" / "owner--repo" / "snapshots" / "v2.0.4" + snapshot.mkdir(parents=True) + (snapshot / "config.json").write_text("{}") + + api = HubApi() + with pytest.raises(CacheNotFound) as exc_info: + api.download_repo( + "owner/repo", + "model", + revision="main", + cache_dir=tmp_path, + local_files_only=True, + ) + + message = str(exc_info.value) + assert "revision 'main'" in message + assert "Cached revisions" in message + assert "v2.0.4" in message diff --git a/tests/test_user_info.py b/tests/test_user_info.py new file mode 100644 index 0000000..daff8c0 --- /dev/null +++ b/tests/test_user_info.py @@ -0,0 +1,103 @@ +"""Tests for tolerant user profile parsing.""" + +from __future__ import annotations + +import pytest + +from modelscope_hub.api import HubApi +from modelscope_hub.types import UserInfo + + +@pytest.mark.parametrize( + ("payload", "expected"), + [ + ( + { + "Username": "legacy-user", + "UserId": 123, + "Email": "legacy@example.com", + "Avatar": "https://avatar.example/legacy.png", + "Description": "legacy description", + }, + UserInfo( + id=123, + username="legacy-user", + email="legacy@example.com", + avatar_url="https://avatar.example/legacy.png", + description="legacy description", + ), + ), + ( + { + "Name": "pre-user", + "sub": "sub-123", + "email": "pre@example.com", + "picture": "https://avatar.example/pre.png", + "description": "pre description", + }, + UserInfo( + id="sub-123", + username="pre-user", + email="pre@example.com", + avatar_url="https://avatar.example/pre.png", + description="pre description", + ), + ), + ( + { + "preferred_username": "display-name", + "name": "login-handle", + "user_id": "uid-1", + "mail": "mail@example.com", + "avatarUrl": "https://avatar.example/a.png", + "bio": "bio text", + }, + UserInfo( + id="uid-1", + username="login-handle", + email="mail@example.com", + avatar_url="https://avatar.example/a.png", + description="bio text", + ), + ), + ({"preferred_username": "fallback-user", "ID": 0}, UserInfo(id=0, username="fallback-user")), + ], +) +def test_user_info_accepts_legacy_and_oidc_field_names(payload, expected): + assert UserInfo.from_dict(payload) == expected + + +def test_user_info_ignores_empty_aliases_until_a_non_empty_value(): + user = UserInfo.from_dict( + { + "Username": "", + "username": None, + "name": "resolved-user", + "Email": "", + "email": "resolved@example.com", + } + ) + + assert user.username == "resolved-user" + assert user.email == "resolved@example.com" + + +def test_whoami_uses_user_info_field_compatibility(monkeypatch): + api = HubApi(token="token-for-test") + monkeypatch.setattr( + api.openapi, + "get_current_user", + lambda: { + "Name": "pre-user", + "sub": "sub-123", + "email": "pre@example.com", + "description": "pre description", + }, + ) + + user = api.whoami() + + assert user.username == "pre-user" + assert user.id == "sub-123" + assert user.email == "pre@example.com" + assert user.description == "pre description"