From 00579c1360ebfec3310f9873d2a445a32870b9b3 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E7=8F=AD=E6=89=AC?= Date: Mon, 20 Jul 2026 10:44:33 +0800 Subject: [PATCH 01/13] bump version --- src/modelscope_hub/version.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/src/modelscope_hub/version.py b/src/modelscope_hub/version.py index cd4cda5..7a71621 100644 --- a/src/modelscope_hub/version.py +++ b/src/modelscope_hub/version.py @@ -1,3 +1,3 @@ """Version information for modelscope_hub.""" -__version__ = "0.1.7+main" +__version__ = "0.1.8+main" From 31f224cc75e6849626ff0c6ac39041236ddc9b9b Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E7=8F=AD=E6=89=AC?= Date: Mon, 20 Jul 2026 11:31:56 +0800 Subject: [PATCH 02/13] fix(download): forward progress_callbacks through HubApi.download_repo facade The compat snapshot_download wrapper forwards progress_callbacks to HubApi.download_repo, but the facade dropped the parameter and never passed it to DownloadManager.download_repo, so any real (non-mocked) call raised TypeError. Restore full-chain propagation: HubApi.download_repo now accepts and forwards progress_callbacks to the DownloadManager delegate (DownloadManager already supports it); compat.snapshot_download accepts and forwards it. Add an integration-style test that mocks only the DownloadManager delegate so the real compat wrapper and HubApi facade both execute. --- src/modelscope_hub/api.py | 7 ++- .../compat/snapshot_download.py | 10 ++++- tests/test_compat_snapshot_download.py | 44 +++++++++++++++++++ 3 files changed, 59 insertions(+), 2 deletions(-) create mode 100644 tests/test_compat_snapshot_download.py diff --git a/src/modelscope_hub/api.py b/src/modelscope_hub/api.py index 0948236..f1642eb 100644 --- a/src/modelscope_hub/api.py +++ b/src/modelscope_hub/api.py @@ -31,7 +31,7 @@ from ._cache_manager import clear_cache as _clear_cache from ._cache_manager import scan_cache as _scan_cache -from ._download import DownloadManager +from ._download import DownloadManager, ProgressCallback from ._legacy_api import LegacyClient from ._openapi import OpenAPIClient from ._upload import UploadManager @@ -1289,6 +1289,7 @@ def download_repo( max_workers: int = 4, local_files_only: bool = False, user_agent: dict | str | None = None, + progress_callbacks: list[type[ProgressCallback]] | None = None, ) -> Path: """Download an entire repository snapshot. @@ -1317,6 +1318,9 @@ def download_repo( When ``True``, return the cached snapshot path without network. user_agent : dict, str or None, optional Custom user-agent info for download headers. + progress_callbacks : list of ProgressCallback subclasses, optional + Callback *classes* (not instances); each is instantiated per file + to report byte-level download progress. Returns ------- @@ -1354,6 +1358,7 @@ def download_repo( max_workers=max_workers, local_files_only=local_files_only, user_agent=user_agent, + progress_callbacks=progress_callbacks, ) def list_repo_files( diff --git a/src/modelscope_hub/compat/snapshot_download.py b/src/modelscope_hub/compat/snapshot_download.py index e021495..00f84a5 100644 --- a/src/modelscope_hub/compat/snapshot_download.py +++ b/src/modelscope_hub/compat/snapshot_download.py @@ -8,7 +8,7 @@ from __future__ import annotations import warnings -from typing import Sequence +from typing import TYPE_CHECKING, Sequence import requests as _requests @@ -19,6 +19,9 @@ from .constants import DEFAULT_DATASET_REVISION from .file_download import _resolve_legacy_paths +if TYPE_CHECKING: + from .._download import ProgressCallback + def snapshot_download( model_id: str | None = None, @@ -38,12 +41,16 @@ def snapshot_download( endpoint: str | None = None, local_files_only: bool = False, user_agent: dict | str | None = None, + progress_callbacks: list[type[ProgressCallback]] | None = None, ) -> str: """Download a repo snapshot (legacy signature). Parameters mirror the old ``modelscope.hub.snapshot_download.snapshot_download``. ``allow_patterns``/``ignore_patterns`` take priority over the ``allow_file_pattern``/``ignore_file_pattern`` aliases when both are set. + ``progress_callbacks`` takes a list of :class:`ProgressCallback` + subclasses (not instances); each is instantiated per file to report + download progress. """ effective_id = repo_id or model_id if not effective_id: @@ -85,6 +92,7 @@ def snapshot_download( max_workers=max_workers, local_files_only=local_files_only, user_agent=user_agent, + progress_callbacks=progress_callbacks, ) except (NotExistError, AuthenticationError, PermissionDeniedError) as e: raise _requests.exceptions.HTTPError( diff --git a/tests/test_compat_snapshot_download.py b/tests/test_compat_snapshot_download.py new file mode 100644 index 0000000..389f1ba --- /dev/null +++ b/tests/test_compat_snapshot_download.py @@ -0,0 +1,44 @@ +"""Unit tests for the ``snapshot_download`` progress_callbacks chain. + +These are network-free integration-style tests: only the *bottom* delegate +(``DownloadManager.download_repo``) is mocked, so the real compat wrapper and +the real ``HubApi.download_repo`` facade both execute. This guarantees the +whole ``compat -> HubApi facade -> DownloadManager`` forwarding chain is +exercised (a facade that drops ``progress_callbacks`` would fail here). +""" +from __future__ import annotations + +from unittest import mock + +from modelscope_hub import ProgressCallback +from modelscope_hub._download import DownloadManager +from modelscope_hub.compat.snapshot_download import snapshot_download + + +class _DummyCallback(ProgressCallback): + pass + + +class TestSnapshotDownloadProgressCallbacks: + def test_progress_callbacks_forwarded_through_facade(self): + with mock.patch.object( + DownloadManager, "download_repo", + return_value="/tmp/snapshot") as m: + result = snapshot_download( + "owner/repo", + progress_callbacks=[_DummyCallback], + local_files_only=True, + ) + + assert str(result) == "/tmp/snapshot" + _, kwargs = m.call_args + assert kwargs["progress_callbacks"] == [_DummyCallback] + + def test_progress_callbacks_default_none(self): + with mock.patch.object( + DownloadManager, "download_repo", + return_value="/tmp/snapshot") as m: + snapshot_download("owner/repo", local_files_only=True) + + _, kwargs = m.call_args + assert kwargs["progress_callbacks"] is None From 306f1459b6e722a51e05967ddca8e9debcc00149 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E7=8F=AD=E6=89=AC?= Date: Tue, 21 Jul 2026 17:04:40 +0800 Subject: [PATCH 03/13] fix(download): harden legacy cache auto-detection for pre-1.38 layouts DownloadManager._find_legacy_repo_dir now probes both historical cache layouts in priority order: {base}/{type}s/{owner}/{name___} (MODELSCOPE_CACHE explicitly set) and {base}/hub/{type}s/{owner}/{name___} (default ~/.cache/modelscope/hub). It returns the first existing non-empty match, or None when the cache is clean so the caller falls back to the new layout. Adds network-free regression tests covering both layouts, multi-dot name encoding, clean cache, and empty legacy dir. --- src/modelscope_hub/_download.py | 25 +++++++++--- tests/test_legacy_cache_detection.py | 58 ++++++++++++++++++++++++++++ 2 files changed, 78 insertions(+), 5 deletions(-) create mode 100644 tests/test_legacy_cache_detection.py diff --git a/src/modelscope_hub/_download.py b/src/modelscope_hub/_download.py index 069556a..4904f83 100644 --- a/src/modelscope_hub/_download.py +++ b/src/modelscope_hub/_download.py @@ -851,8 +851,18 @@ def _find_legacy_repo_dir( ) -> Path | None: """Check for old SDK (<=1.37) cache layout and return it if non-empty. - Old format: {base}/{type}s/{owner}/{name_with_dots_as___}/ - e.g. ~/.cache/modelscope/models/Qwen/Qwen3___5-0___8B/ + Old SDKs encoded dots in the repo name as ``___`` and, when no + ``MODELSCOPE_CACHE`` was set, stored everything under a ``hub/`` + sub-directory of the cache root. Both historical layouts are probed, + in priority order:: + + {base}/{type}s/{owner}/{name___} # MODELSCOPE_CACHE explicitly set + {base}/hub/{type}s/{owner}/{name___} # default cache (~/.cache/modelscope/hub) + + e.g. ~/.cache/modelscope/hub/models/Qwen/Qwen3___5-0___8B/ + + 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 segment = f"{repo_type}s" if not repo_type.endswith("s") else repo_type @@ -861,13 +871,18 @@ def _find_legacy_repo_dir( return None owner, name = parts safe_name = name.replace(".", "___") - legacy_path = base / segment / owner / safe_name - if legacy_path.is_dir(): + candidates = ( + base / segment / owner / safe_name, + base / "hub" / segment / owner / safe_name, + ) + for legacy_path in candidates: + if not legacy_path.is_dir(): + continue try: if any(legacy_path.iterdir()): return legacy_path except OSError: - pass + continue return None def _lock_path( diff --git a/tests/test_legacy_cache_detection.py b/tests/test_legacy_cache_detection.py new file mode 100644 index 0000000..e7c46f3 --- /dev/null +++ b/tests/test_legacy_cache_detection.py @@ -0,0 +1,58 @@ +"""Unit tests for legacy (pre-1.38) cache auto-detection. + +These are network-free tests for ``DownloadManager._find_legacy_repo_dir``, +which lets ``download_repo`` / ``download_file`` reuse an existing old-SDK +cache instead of re-downloading into the new layout. +""" +from __future__ import annotations + +from modelscope_hub.api import HubApi + + +def _make_download_manager(): + """Build a network-free DownloadManager via the public HubApi facade.""" + return HubApi().downloader + + +class TestFindLegacyRepoDir: + def test_detects_modelscope_cache_layout(self, tmp_path): + # MODELSCOPE_CACHE explicitly set: {base}/models/{owner}/{name___} + legacy = tmp_path / "models" / "Qwen" / "Qwen3___5-4B" + legacy.mkdir(parents=True) + (legacy / "config.json").write_text("{}") + + dm = _make_download_manager() + found = dm._find_legacy_repo_dir("Qwen/Qwen3.5-4B", "model", tmp_path) + assert found == legacy + + def test_detects_default_hub_segment_layout(self, tmp_path): + # Default cache (~/.cache/modelscope/hub): {base}/hub/models/{owner}/{name___} + legacy = tmp_path / "hub" / "models" / "Qwen" / "Qwen3___5-4B" + legacy.mkdir(parents=True) + (legacy / "config.json").write_text("{}") + + dm = _make_download_manager() + found = dm._find_legacy_repo_dir("Qwen/Qwen3.5-4B", "model", tmp_path) + assert found == legacy + + def test_multi_dot_name_encoding(self, tmp_path): + legacy = tmp_path / "models" / "Qwen" / "Qwen2___5-0___5B" + legacy.mkdir(parents=True) + (legacy / "config.json").write_text("{}") + + dm = _make_download_manager() + found = dm._find_legacy_repo_dir("Qwen/Qwen2.5-0.5B", "model", tmp_path) + assert found == legacy + + def test_clean_cache_returns_none(self, tmp_path): + dm = _make_download_manager() + assert dm._find_legacy_repo_dir( + "Qwen/Qwen3.5-4B", "model", tmp_path) is None + + def test_empty_legacy_dir_returns_none(self, tmp_path): + legacy = tmp_path / "models" / "Qwen" / "Qwen3___5-4B" + legacy.mkdir(parents=True) # exists but empty + + dm = _make_download_manager() + assert dm._find_legacy_repo_dir( + "Qwen/Qwen3.5-4B", "model", tmp_path) is None From 07aec55e0a20a21520812ff1eb0bf6202b79a2b1 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E7=8F=AD=E6=89=AC?= Date: Tue, 21 Jul 2026 18:28:48 +0800 Subject: [PATCH 04/13] fix(packaging): rename console scripts to modelscope-hub/ms-hub to avoid conflict modelscope and modelscope-hub both declared modelscope/ms console scripts, so installing both (modelscope depends on modelscope-hub) collided on the same binary (e.g. /usr/local/bin/modelscope-3.12 on FreeBSD pkg). Rename the hub entry points to modelscope-hub/ms-hub so the umbrella modelscope package keeps modelscope/ms while standalone modelscope-hub users get their own non-colliding CLI. Fixes modelscope/modelscope#1752. --- pyproject.toml | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/pyproject.toml b/pyproject.toml index 81b0018..898b356 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -35,8 +35,8 @@ dev = [ ] [project.scripts] -modelscope = "modelscope_hub.cli.main:run_cmd" -ms = "modelscope_hub.cli.main:run_cmd" +modelscope-hub = "modelscope_hub.cli.main:run_cmd" +ms-hub = "modelscope_hub.cli.main:run_cmd" [build-system] requires = ["setuptools>=68.0"] From 47b866a76a76919932e3c48f58866bf3f2c2fb1f Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E7=8F=AD=E6=89=AC?= Date: Tue, 21 Jul 2026 18:49:38 +0800 Subject: [PATCH 05/13] update cli: ms/modelscope -> ms-hub/modelscope-hub --- README.md | 198 ++++++++++++++++----------------- src/modelscope_hub/cli/main.py | 18 ++- 2 files changed, 113 insertions(+), 103 deletions(-) diff --git a/README.md b/README.md index 4cdd7a7..a2aee78 100644 --- a/README.md +++ b/README.md @@ -21,7 +21,7 @@ The official Python SDK & CLI for [ModelScope Hub](https://modelscope.cn) — do ## Why modelscope-hub? -`modelscope-hub` connects your code to the [ModelScope](https://modelscope.cn) ecosystem — models, datasets, Studio spaces, skills, and MCP servers — through a single `HubApi` class or the `ms` CLI. +`modelscope-hub` connects your code to the [ModelScope](https://modelscope.cn) ecosystem — models, datasets, Studio spaces, skills, and MCP servers — through a single `HubApi` class or the `ms-hub` CLI. - **Unified repo interface** — one set of methods for models, datasets, studios, skills, and MCP servers - **OpenAPI-first** — built on the ModelScope OpenAPI surface with transparent legacy fallback @@ -46,7 +46,7 @@ The official Python SDK & CLI for [ModelScope Hub](https://modelscope.cn) — do - **Fix**: `parse_timestamp` robust timezone conversion for ISO 8601, floats, milliseconds **v0.0.8** (2026-06-10) -- **Feature**: `ms list --all` auto-pagination; `ms create --skill-file` zip upload; `ms list --envs` +- **Feature**: `ms-hub list --all` auto-pagination; `ms-hub create --skill-file` zip upload; `ms-hub list --envs` - **Fix**: Download per-file lock & stale detection & atomic merge; `--disable-tqdm` for folder upload - **Security**: Redact tokens from git/API error output - **Refactor**: Centralize env var registry; unify `MODELSCOPE_DOMAIN` → `MODELSCOPE_ENDPOINT` @@ -59,7 +59,7 @@ The official Python SDK & CLI for [ModelScope Hub](https://modelscope.cn) — do - OpenAPI spec alignment: pagination limits, retry, auth, request body **v0.0.4** (2026-06-05) -- Flatten CLI to top-level commands (`ms create/info/list/delete`) +- Flatten CLI to top-level commands (`ms-hub create/info/list/delete`) - Migrate credentials to `~/.modelscope/credentials/` - Fix dataset/skill download, blob upload auth, error code refactor @@ -82,9 +82,9 @@ Requires Python 3.10+. Lightweight — only `requests`, `tqdm`, `filelock`, `url ### Authenticate ```bash -ms login +ms-hub login # or pass a token directly -ms login --token $MODELSCOPE_API_TOKEN +ms-hub login --token $MODELSCOPE_API_TOKEN ``` Get your token at [modelscope.cn/my/access/token](https://modelscope.cn/my/access/token) or [modelscope.ai/my/access/token](https://modelscope.ai/my/access/token). @@ -101,16 +101,16 @@ print(user.username) ```bash # Full snapshot -ms download Qwen/Qwen3-0.6B +ms-hub download Qwen/Qwen3-0.6B # Single file -ms download Qwen/Qwen3-0.6B config.json +ms-hub download Qwen/Qwen3-0.6B config.json # With filters -ms download Qwen/Qwen3-0.6B --include "*.safetensors" --exclude "*.bin" +ms-hub download Qwen/Qwen3-0.6B --include "*.safetensors" --exclude "*.bin" # Directly into a local directory (bypasses cache) -ms download Qwen/Qwen3-0.6B --local-dir ./my-model +ms-hub download Qwen/Qwen3-0.6B --local-dir ./my-model ``` ```python @@ -129,8 +129,8 @@ path = api.download_file("Qwen/Qwen3-0.6B", "model", "config.json", local_files_ ### Upload ```bash -ms upload my-org/my-model ./weights.safetensors -ms upload my-org/my-model ./output --repo-type model --commit-message "add weights" +ms-hub upload my-org/my-model ./weights.safetensors +ms-hub upload my-org/my-model ./output --repo-type model --commit-message "add weights" ``` ```python @@ -141,7 +141,7 @@ api.upload_folder("my-org/my-model", "model", "./output", path_in_repo="") ### Create a Repository ```bash -ms create my-org/my-model --repo-type model --visibility private +ms-hub create my-org/my-model --repo-type model --visibility private ``` ```python @@ -151,9 +151,9 @@ api.create_repo("my-org/my-model", "model", visibility="private", license="apach ### Deploy a Studio ```bash -ms deploy my-org/chat-demo --repo-type studio -ms logs my-org/chat-demo --log-type run -ms stop my-org/chat-demo --repo-type studio +ms-hub deploy my-org/chat-demo --repo-type studio +ms-hub logs my-org/chat-demo --log-type run +ms-hub stop my-org/chat-demo --repo-type studio ``` ```python @@ -166,7 +166,7 @@ api.stop_repo("my-org/chat-demo", "studio") ## CLI Reference -The CLI is available as both `ms` and `modelscope`. +The CLI is available as both `ms-hub` and `modelscope-hub`. **Global options** (placed before or after the subcommand): @@ -178,40 +178,40 @@ The CLI is available as both `ms` and `modelscope`. | `-V, --version` | Print version and exit (global only) | > `--token` and `--endpoint` can be placed either before or after the subcommand: -> `ms --token xxx download ...` and `ms download ... --token xxx` are equivalent. +> `ms-hub --token xxx download ...` and `ms-hub download ... --token xxx` are equivalent. -### `ms login` +### `ms-hub login` Authenticate and persist your token locally. ```bash -ms login # interactive prompt -ms login --token $MY_TOKEN # non-interactive +ms-hub login # interactive prompt +ms-hub login --token $MY_TOKEN # non-interactive ``` | Option | Description | |--------|-------------| | `--token TOKEN` | API token; prompted interactively if omitted | -### `ms whoami` +### `ms-hub whoami` Show the user associated with the current token. ```bash -ms whoami -ms whoami --token $MY_TOKEN # check a specific token without logging in +ms-hub whoami +ms-hub whoami --token $MY_TOKEN # check a specific token without logging in ``` -### `ms download` +### `ms-hub download` Download a single file or a full repository snapshot. ```bash -ms download Qwen/Qwen3-0.6B # full snapshot -ms download Qwen/Qwen3-0.6B config.json # single file -ms download Qwen/Qwen3-0.6B --include "*.safetensors" # filter by glob -ms download Qwen/Qwen3-0.6B --local-dir ./out --max-workers 8 # direct download -ms download my-org/my-data --repo-type dataset --revision v2 # dataset at tag +ms-hub download Qwen/Qwen3-0.6B # full snapshot +ms-hub download Qwen/Qwen3-0.6B config.json # single file +ms-hub download Qwen/Qwen3-0.6B --include "*.safetensors" # filter by glob +ms-hub download Qwen/Qwen3-0.6B --local-dir ./out --max-workers 8 # direct download +ms-hub download my-org/my-data --repo-type dataset --revision v2 # dataset at tag ``` | Argument / Option | Required | Description | @@ -232,40 +232,40 @@ ms download my-org/my-data --repo-type dataset --revision v2 # dataset at tag ```bash # Download multiple specific files at once -ms download Qwen/Qwen3-0.6B config.json tokenizer.json generation_config.json +ms-hub download Qwen/Qwen3-0.6B config.json tokenizer.json generation_config.json # Download only safetensors, skip GGUF and bin weights -ms download Qwen/Qwen3-0.6B --include "*.safetensors" --exclude "*.bin" "*.gguf" +ms-hub download Qwen/Qwen3-0.6B --include "*.safetensors" --exclude "*.bin" "*.gguf" # Download a dataset at a specific tag into a local directory -ms download my-org/my-data --repo-type dataset --revision v2 --local-dir ./data +ms-hub download my-org/my-data --repo-type dataset --revision v2 --local-dir ./data # Use a custom cache directory and 8 parallel threads -ms download Qwen/Qwen3-0.6B --cache-dir /data/hub-cache --max-workers 8 +ms-hub download Qwen/Qwen3-0.6B --cache-dir /data/hub-cache --max-workers 8 # Force re-download even if already cached -ms download Qwen/Qwen3-0.6B config.json --force +ms-hub download Qwen/Qwen3-0.6B config.json --force # Download all skills from a collection (legacy flag) -ms download --collection my-org/skill-collection +ms-hub download --collection my-org/skill-collection # Enable parallel range download for large files (env var) -MODELSCOPE_DOWNLOAD_PARALLELS=4 ms download Qwen/Qwen3-0.6B +MODELSCOPE_DOWNLOAD_PARALLELS=4 ms-hub download Qwen/Qwen3-0.6B # Use the modelscope.ai endpoint (global option, before subcommand) -ms --endpoint https://modelscope.ai download Qwen/Qwen3-0.6B +ms-hub --endpoint https://modelscope.ai download Qwen/Qwen3-0.6B ``` -### `ms upload` +### `ms-hub upload` Upload a file or folder to a repository. ```bash -ms upload my-org/my-model ./weights.safetensors # single file -ms upload my-org/my-model ./output models/ --repo-type model # folder → subdir -ms upload my-org/my-model . --include "*.py" --commit-message "code" # filtered folder +ms-hub upload my-org/my-model ./weights.safetensors # single file +ms-hub upload my-org/my-model ./output models/ --repo-type model # folder → subdir +ms-hub upload my-org/my-model . --include "*.py" --commit-message "code" # filtered folder ``` | Argument / Option | Required | Description | @@ -288,56 +288,56 @@ ms upload my-org/my-model . --include "*.py" --commit-message "code" # filtered ```bash # Upload a single file with a custom commit message -ms upload my-org/my-model ./weights.safetensors --commit-message "add fp16 weights" +ms-hub upload my-org/my-model ./weights.safetensors --commit-message "add fp16 weights" # Upload a folder into a subdirectory of the repo -ms upload my-org/my-model ./output models/ --repo-type model +ms-hub upload my-org/my-model ./output models/ --repo-type model # Upload only Python files from the current directory -ms upload my-org/my-model . --include "*.py" --commit-message "update code" +ms-hub upload my-org/my-model . --include "*.py" --commit-message "update code" # Upload only safetensors, skip checkpoints -ms upload my-org/my-model ./output --include "*.safetensors" --exclude "*.ckpt" "*.bin" +ms-hub upload my-org/my-model ./output --include "*.safetensors" --exclude "*.ckpt" "*.bin" # Upload to a dataset repo on a specific branch -ms upload my-org/my-data ./data --repo-type dataset --revision dev +ms-hub upload my-org/my-data ./data --repo-type dataset --revision dev # Upload with extended commit description -ms upload my-org/my-model ./weights.safetensors \ +ms-hub upload my-org/my-model ./weights.safetensors \ --commit-message "v2 weights" \ --commit-description "Retrained with extended dataset, 3 epochs, lr=2e-5" # Resumable upload: interrupted uploads resume automatically via cache -ms upload my-org/my-model ./large-folder +ms-hub upload my-org/my-model ./large-folder # If interrupted, just re-run the same command — already uploaded files are skipped # Disable upload cache (no resume, fresh upload every time) -ms upload my-org/my-model ./output --no-cache +ms-hub upload my-org/my-model ./output --no-cache # Disable progress bars (useful for CI/CD pipelines) -ms upload my-org/my-model ./output --disable-tqdm +ms-hub upload my-org/my-model ./output --disable-tqdm ``` -### `ms create` / `ms info` / `ms list` / `ms delete` +### `ms-hub create` / `ms-hub info` / `ms-hub list` / `ms-hub delete` Repository management. ```bash -ms create my-org/my-model --repo-type model --visibility private -ms create my-org/demo --repo-type studio --sdk-type gradio -ms info my-org/my-model --repo-type model -ms list --repo-type model --owner my-org --page-size 20 -ms delete my-org/my-model --repo-type model --yes +ms-hub create my-org/my-model --repo-type model --visibility private +ms-hub create my-org/demo --repo-type studio --sdk-type gradio +ms-hub info my-org/my-model --repo-type model +ms-hub list --repo-type model --owner my-org --page-size 20 +ms-hub delete my-org/my-model --repo-type model --yes ``` -> **Deprecation notice:** `delete_repo` / `ms delete` emits a `DeprecationWarning` — programmatic repo deletion is restricted for security reasons and will be restored once token-scoped auth is available. Use the [web console](https://modelscope.cn) to delete repos. +> **Deprecation notice:** `delete_repo` / `ms-hub delete` emits a `DeprecationWarning` — programmatic repo deletion is restricted for security reasons and will be restored once token-scoped auth is available. Use the [web console](https://modelscope.cn) to delete repos. > > `delete_files` requires cookie-based session auth; API tokens may receive a 401 error.
-ms create options +ms-hub create options | Argument / Option | Required | Description | |-------------------|----------|-------------| @@ -356,15 +356,15 @@ ms delete my-org/my-model --repo-type model --yes
-### `ms deploy` / `ms stop` / `ms logs` / `ms settings` +### `ms-hub deploy` / `ms-hub stop` / `ms-hub logs` / `ms-hub settings` Manage Studio and MCP deployments. ```bash -ms deploy my-org/chat-demo --repo-type studio -ms logs my-org/chat-demo --log-type run --keyword ERROR --page-size 50 -ms settings my-org/chat-demo cpu=4 memory=8192 -ms stop my-org/chat-demo --repo-type studio +ms-hub deploy my-org/chat-demo --repo-type studio +ms-hub logs my-org/chat-demo --log-type run --keyword ERROR --page-size 50 +ms-hub settings my-org/chat-demo cpu=4 memory=8192 +ms-hub stop my-org/chat-demo --repo-type studio ```
@@ -372,25 +372,25 @@ ms stop my-org/chat-demo --repo-type studio | Command | `--repo-type` | Key Options | |---------|---------------|-------------| -| `ms deploy ` | `{studio,mcp}` (default: `studio`) | — | -| `ms stop ` | `{studio,mcp}` (default: `studio`) | — | -| `ms logs ` | `{studio}` only | `--log-type {run,build}`, `--keyword`, `--page`, `--page-size` | -| `ms settings key=val...` | `{studio,skill}` (default: `studio`) | Key-value pairs passed to backend | +| `ms-hub deploy ` | `{studio,mcp}` (default: `studio`) | — | +| `ms-hub stop ` | `{studio,mcp}` (default: `studio`) | — | +| `ms-hub logs ` | `{studio}` only | `--log-type {run,build}`, `--keyword`, `--page`, `--page-size` | +| `ms-hub settings key=val...` | `{studio,skill}` (default: `studio`) | Key-value pairs passed to backend | -> **Note:** `ms logs` only supports Studio spaces. MCP server logs are not available via this command. -> `ms settings` supports Studio and Skill repos; for MCP servers use `ms mcp deploy` with configuration payload. +> **Note:** `ms-hub logs` only supports Studio spaces. MCP server logs are not available via this command. +> `ms-hub settings` supports Studio and Skill repos; for MCP servers use `ms-hub mcp deploy` with configuration payload.
-### `ms secret` +### `ms-hub secret` Manage secrets for Studio spaces (studio only, `--repo-type` defaults to `studio`). ```bash -ms secret add my-org/demo API_KEY sk-xxx -ms secret list my-org/demo -ms secret update my-org/demo API_KEY sk-new -ms secret delete my-org/demo API_KEY --yes +ms-hub secret add my-org/demo API_KEY sk-xxx +ms-hub secret list my-org/demo +ms-hub secret update my-org/demo API_KEY sk-new +ms-hub secret delete my-org/demo API_KEY --yes ```
@@ -407,15 +407,15 @@ All subcommands accept `--repo-type` (default: `studio`, currently the only supp
-### `ms mcp` +### `ms-hub mcp` Manage MCP (Model Context Protocol) servers. ```bash -ms mcp list --search weather --page-size 10 -ms mcp info my-org/weather-mcp -ms mcp deploy my-org/weather-mcp -ms mcp undeploy my-org/weather-mcp +ms-hub mcp list --search weather --page-size 10 +ms-hub mcp info my-org/weather-mcp +ms-hub mcp deploy my-org/weather-mcp +ms-hub mcp undeploy my-org/weather-mcp ```
@@ -430,17 +430,17 @@ ms mcp undeploy my-org/weather-mcp
-### `ms cache` +### `ms-hub cache` Inspect and clean the local download cache. ```bash -ms cache scan -ms cache scan --cache-dir /data/cache -ms cache verify Qwen/Qwen3-0.6B -ms cache verify Qwen/Qwen3-0.6B --local-dir ./Qwen3-0.6B -ms cache clear --repo-type model --yes -ms cache clear --repo-id my-org/old-model --repo-type model --yes +ms-hub cache scan +ms-hub cache scan --cache-dir /data/cache +ms-hub cache verify Qwen/Qwen3-0.6B +ms-hub cache verify Qwen/Qwen3-0.6B --local-dir ./Qwen3-0.6B +ms-hub cache clear --repo-type model --yes +ms-hub cache clear --repo-id my-org/old-model --repo-type model --yes ```
@@ -454,13 +454,13 @@ ms cache clear --repo-id my-org/old-model --repo-type model --yes
-### `ms agent` +### `ms-hub agent` Low-level raw file transfer for remote agent repositories: `download`, `upload`, `list`. This command transfers files as-is, with **no framework awareness**. ```bash -ms agent download -r user/my-agent --local-dir ./my-agent # download raw files -ms agent upload -r user/my-agent --local-dir ./my-agent # upload raw +ms-hub agent download -r user/my-agent --local-dir ./my-agent # download raw files +ms-hub agent upload -r user/my-agent --local-dir ./my-agent # upload raw ``` > **Framework-aware operations** (cross-framework `convert`, `watch`/bidirectional sync, `status`, `backups`, `restore`, `stop`) live in **[modelscope-agent](https://github.com/modelscope/ms-agent)** — use `ms-agent agent ...` instead. For example, to download and convert in one step: `ms-agent agent download -f qoder -r user/my-agent --target-framework qwenpaw`. @@ -468,13 +468,13 @@ ms agent upload -r user/my-agent --local-dir ./my-agent # upload raw
Subcommands -#### `ms agent download` +#### `ms-hub agent download` Download all files of a remote agent repository to a local directory (raw, no conversion). ```bash -ms agent download -r user/my-agent -ms agent download -r user/my-agent --local-dir ./my-agent --revision master +ms-hub agent download -r user/my-agent +ms-hub agent download -r user/my-agent --local-dir ./my-agent --revision master ``` | Option | Required | Description | @@ -483,13 +483,13 @@ ms agent download -r user/my-agent --local-dir ./my-agent --revision master | `--local-dir DIR` | no | Destination directory (default: `./` under CWD) | | `--revision REV` | no | Repository revision (default: `master`) | -#### `ms agent upload` +#### `ms-hub agent upload` Upload files from a local path (file or directory) to a remote agent repository (raw, no conversion). Creates the repo if it does not exist. ```bash -ms agent upload -r user/my-agent --local-dir ./my-agent -ms agent upload -r user/my-agent --local-dir ./my-agent --dry-run +ms-hub agent upload -r user/my-agent --local-dir ./my-agent +ms-hub agent upload -r user/my-agent --local-dir ./my-agent --dry-run ``` | Option | Required | Description | @@ -585,7 +585,7 @@ api = HubApi(token="...", endpoint="https://modelscope.ai") (training · eval) (inference · deploy) ``` -- **Browse & discover** — search 100K+ models and datasets via `list_repos` / `ms repo list` +- **Browse & discover** — search 100K+ models and datasets via `list_repos` / `ms-hub repo list` - **Download & cache** — pull model weights, tokenizer configs, or entire datasets into a managed cache or a local directory; supports offline mode via `local_files_only` - **Train & fine-tune** — use with the [modelscope](https://github.com/modelscope/modelscope) framework: train locally, then push results back - **Deploy** — launch a Studio space or MCP server directly from the CLI or SDK @@ -595,9 +595,9 @@ api = HubApi(token="...", endpoint="https://modelscope.ai") ## Configuration -Run `ms list --envs` to see all configurable environment variables with their current values. +Run `ms-hub list --envs` to see all configurable environment variables with their current values. -Token is persisted locally after `ms login` and auto-loaded in subsequent sessions. +Token is persisted locally after `ms-hub login` and auto-loaded in subsequent sessions.
Environment variables @@ -654,7 +654,7 @@ Token is persisted locally after `ms login` and auto-loaded in subsequent sessio | `MODELSCOPE_NO_DEPRECATION_WARNINGS` | — | Suppress deprecation warnings | > Old variable names (e.g. `API_TIMEOUT`, `DOWNLOAD_RETRY_TIMES`, `UPLOAD_USE_CACHE`) are -> still accepted but emit a `FutureWarning`. Run `ms list --envs` to see which deprecated +> still accepted but emit a `FutureWarning`. Run `ms-hub list --envs` to see which deprecated > names are active in your environment.
diff --git a/src/modelscope_hub/cli/main.py b/src/modelscope_hub/cli/main.py index 096c2d4..f46a7e4 100644 --- a/src/modelscope_hub/cli/main.py +++ b/src/modelscope_hub/cli/main.py @@ -1,4 +1,10 @@ -"""Entry point for the ``modelscope`` / ``ms`` console scripts. +"""Entry point for the ``modelscope-hub`` / ``ms-hub`` console scripts. + +This module is also reused by the umbrella ``modelscope`` package for its +``modelscope`` / ``ms`` commands. Because a single parser serves all four +aliases, the program name is derived from ``sys.argv[0]`` (rather than +hard-coded) so help/usage output shows whichever command was actually +invoked. Subcommands live in dedicated modules and are wired in via their :meth:`CLICommand.register` static method. :func:`run_cmd` is intentionally @@ -55,8 +61,12 @@ def _build_parser() -> argparse.ArgumentParser: + # ``prog`` is intentionally left unset so argparse derives it from + # ``sys.argv[0]``. The same parser backs the standalone + # ``modelscope-hub`` / ``ms-hub`` scripts and the umbrella + # ``modelscope`` / ``ms`` scripts, so help output reflects whichever + # command the user actually ran. parser = argparse.ArgumentParser( - prog="ms", description="ModelScope Hub command-line interface.", ) parser.add_argument( @@ -107,14 +117,14 @@ def _register_aliases(subparsers) -> None: def _register_scan_cache_alias(subparsers) -> None: - """``ms scan-cache`` → alias for ``ms cache scan``.""" + """``ms-hub scan-cache`` → alias for ``ms-hub cache scan``.""" p = subparsers.add_parser("scan-cache", help="[Alias] Show cached repos and disk usage.") p.add_argument("--dir", "--cache-dir", dest="cache_dir", default=None) p.set_defaults(_command=_ScanCacheAlias) def _register_clear_cache_alias(subparsers) -> None: - """``ms clear-cache`` → alias for ``ms cache clear``.""" + """``ms-hub clear-cache`` → alias for ``ms-hub cache clear``.""" from ..constants import RepoType p = subparsers.add_parser("clear-cache", help="[Alias] Remove cached files.") From dde187575d747f27f19aaeb509b786d57ea10468 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E7=8F=AD=E6=89=AC?= Date: Tue, 21 Jul 2026 19:03:25 +0800 Subject: [PATCH 06/13] docs(readme): expand recent version news, fold older, group by type Bring the News section up to date through v0.1.8 (v0.1.5-v0.1.8 expanded, v0.1.4 and older collapsed under a details block). Within each release, merge same-type entries (Feature/Fix) into a single bullet instead of listing them flat. CLI references use the new ms-hub command name. --- README.md | 40 +++++++++++++++++++++++++++++++++++----- 1 file changed, 35 insertions(+), 5 deletions(-) diff --git a/README.md b/README.md index a2aee78..436278e 100644 --- a/README.md +++ b/README.md @@ -33,6 +33,40 @@ The official Python SDK & CLI for [ModelScope Hub](https://modelscope.cn) — do ## News +**v0.1.8** (2026-07-21) +- **Feature**: `ms-hub agent` raw file transfer (download/upload/list) for remote agent repos; visibility support for agent hub; cache checksum verification (`ms-hub cache verify`) +- **Fix**: forward `progress_callbacks` through `HubApi.download_repo` so custom download-progress callbacks work end-to-end; harden legacy (pre-1.38) cache auto-detection (reuse existing `{cache}/models/...` and default `{cache}/hub/models/...` layouts); normal (non-LFS) file upload +- **Packaging**: rename console scripts to `modelscope-hub` / `ms-hub` to avoid a file conflict with the `modelscope` package (e.g. FreeBSD pkg) + +**v0.1.7** (2026-07-07) +- **Feature**: intra-/inter-region cloud download acceleration, with a source marker in the progress bar +- **Fix**: align `snapshot_download` cache path with the CLI; add legacy cache fallback +- **Refactor**: inter-region config via env var only (removed the `--inter-regions` CLI arg); cache the region probe + +**v0.1.6** (2026-07-03) +- **Refactor**: replace the extra-field whitelist with a reserved-field blocklist for more permissive param passthrough + +**v0.1.5** (2026-06-30) +- **Fix**: adaptive commit batch size for uploads + +
+Older releases + +**v0.1.4** (2026-06-26) +- **Feature**: `gated_mode` parameter for `create_repo`; `ms-hub create --gated/--no-gated` flags +- **Refactor**: unify visibility / gated_mode semantics in the SDK layer +- **Fix**: `create_repo` extra-kwargs whitelist + type validation; correct visibility mapping (`private` bool is authoritative) + +**v0.1.3** (2026-06-23) +- **Feature**: add `AlreadyExistsError` (E3026) and fix the `exist_ok` mechanism; align `list_repos`/`RepoInfo` with the OpenAPI response format +- **Fix**: `clear-cache` supports all cache layouts (standard/flat/legacy); add `last_modified` mapping and `to_dict()` for `RepoInfo`/`PagedResult` + +**v0.1.2** (2026-06-23) +- **Fix**: unify `list_datasets`/`get_dataset` return format and align parameters + +**v0.1.1** (2026-06-22) +- **Fix**: legacy API for msdatasets loading + **v0.1.0** (2026-06-18) - **Feature**: Configurable upload failure thresholds (consecutive failures & total wait time) - **Fix**: compatibility && error handling @@ -42,8 +76,7 @@ The official Python SDK & CLI for [ModelScope Hub](https://modelscope.cn) — do **v0.0.9** (2026-06-12) - **Feature**: `get_model` support `revision`; expanded param passthrough for repo/model ops -- **Fix**: Pattern normalization accepts iterable inputs (tuple, etc.) -- **Fix**: `parse_timestamp` robust timezone conversion for ISO 8601, floats, milliseconds +- **Fix**: Pattern normalization accepts iterable inputs (tuple, etc.); `parse_timestamp` robust timezone conversion for ISO 8601, floats, milliseconds **v0.0.8** (2026-06-10) - **Feature**: `ms-hub list --all` auto-pagination; `ms-hub create --skill-file` zip upload; `ms-hub list --envs` @@ -51,9 +84,6 @@ The official Python SDK & CLI for [ModelScope Hub](https://modelscope.cn) — do - **Security**: Redact tokens from git/API error output - **Refactor**: Centralize env var registry; unify `MODELSCOPE_DOMAIN` → `MODELSCOPE_ENDPOINT` -
-Older releases - **v0.0.5** (2026-06-05) - Fix `list_repos` pagination and dataset visibility issues - OpenAPI spec alignment: pagination limits, retry, auth, request body From 36659786b59527e3be06c5c28668161075c8cf25 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E7=8F=AD=E6=89=AC?= Date: Wed, 22 Jul 2026 10:48:39 +0800 Subject: [PATCH 07/13] fix revision pass --- src/modelscope_hub/compat/hub_api.py | 31 +++++++++++-- tests/test_compat_get_model_files.py | 67 ++++++++++++++++++++++++++++ 2 files changed, 94 insertions(+), 4 deletions(-) create mode 100644 tests/test_compat_get_model_files.py diff --git a/src/modelscope_hub/compat/hub_api.py b/src/modelscope_hub/compat/hub_api.py index 9ae9730..2fbb44b 100644 --- a/src/modelscope_hub/compat/hub_api.py +++ b/src/modelscope_hub/compat/hub_api.py @@ -96,10 +96,33 @@ def get_model(self, model_id: str, revision: str | None = None) -> dict: info = self._api.get_repo(model_id, RepoType.MODEL, revision=revision) return _repo_info_to_dict(info) - def get_model_files(self, model_id: str, recursive: bool = True) -> list[dict]: - """List files in a model repo.""" - files = self._api.list_repo_files(model_id, RepoType.MODEL, recursive=recursive) - return [{"Path": f.path, "Size": f.size} for f in files] + def get_model_files( + self, + model_id: str, + revision: str | None = None, + root: str | None = None, + recursive: bool = True, + **kwargs: Any, + ) -> list[dict]: + """List files in a model repo. + + Legacy-compatible signature: ``revision`` selects a branch/tag/commit, + ``root`` restricts results to files under a sub-path, and ``recursive`` + walks subdirectories. Other legacy transport kwargs (e.g. + ``use_cookies``, ``headers``) are accepted and ignored — auth is + handled by the configured token/session. + """ + files = self._api.list_repo_files( + model_id, RepoType.MODEL, revision=revision, recursive=recursive, + ) + result = [{"Path": f.path, "Size": f.size} for f in files] + if root: + prefix = root.strip("/") + result = [ + f for f in result + if f["Path"] == prefix or f["Path"].startswith(prefix + "/") + ] + return result def create_repo( self, diff --git a/tests/test_compat_get_model_files.py b/tests/test_compat_get_model_files.py new file mode 100644 index 0000000..1ab39dc --- /dev/null +++ b/tests/test_compat_get_model_files.py @@ -0,0 +1,67 @@ +"""Unit tests for the legacy-compatible ``LegacyHubApi.get_model_files``. + +Network-free: the underlying ``HubApi.list_repo_files`` is mocked so we only +verify the compat wrapper's signature and parameter forwarding. Regression +guard for callers (e.g. vLLM) that pass the historical ``revision`` / ``root`` +keyword arguments. +""" +from __future__ import annotations + +from types import SimpleNamespace +from unittest import mock + +from modelscope_hub.compat import LegacyHubApi + + +def _fake_files(): + return [ + SimpleNamespace(path="config.json", size=10), + SimpleNamespace(path="model.safetensors", size=100), + SimpleNamespace(path="subdir/extra.bin", size=5), + ] + + +class TestGetModelFilesLegacyCompat: + def test_revision_is_accepted_and_forwarded(self): + lha = LegacyHubApi() + with mock.patch.object( + lha._api, "list_repo_files", + return_value=_fake_files()) as m: + out = lha.get_model_files( + "Qwen/Qwen2.5-1.5B-Instruct", revision="v2") + + assert [f["Path"] for f in out] == [ + "config.json", "model.safetensors", "subdir/extra.bin", + ] + _, kwargs = m.call_args + assert kwargs["revision"] == "v2" + + def test_root_restricts_to_subpath(self): + lha = LegacyHubApi() + with mock.patch.object( + lha._api, "list_repo_files", return_value=_fake_files()): + out = lha.get_model_files("owner/name", root="subdir") + + assert [f["Path"] for f in out] == ["subdir/extra.bin"] + + def test_tolerates_legacy_transport_kwargs(self): + lha = LegacyHubApi() + with mock.patch.object( + lha._api, "list_repo_files", return_value=_fake_files()): + # Historical kwargs must not raise "unexpected keyword argument". + out = lha.get_model_files( + "owner/name", revision="master", + use_cookies=True, headers={}) + + assert len(out) == 3 + + def test_default_revision_none_forwarded(self): + lha = LegacyHubApi() + with mock.patch.object( + lha._api, "list_repo_files", + return_value=_fake_files()) as m: + lha.get_model_files("owner/name") + + _, kwargs = m.call_args + assert kwargs["revision"] is None + assert kwargs["recursive"] is True From 17244b2e0f5c0df0730e3c3777073beac8964311 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E7=8F=AD=E6=89=AC?= Date: Fri, 31 Jul 2026 15:06:55 +0800 Subject: [PATCH 08/13] bump version --- src/modelscope_hub/version.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/src/modelscope_hub/version.py b/src/modelscope_hub/version.py index 7a71621..8804a1a 100644 --- a/src/modelscope_hub/version.py +++ b/src/modelscope_hub/version.py @@ -1,3 +1,3 @@ """Version information for modelscope_hub.""" -__version__ = "0.1.8+main" +__version__ = "0.1.9+main" From 91a54895f7fab07e5a43433e2848258409067d71 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E7=8F=AD=E6=89=AC?= Date: Fri, 31 Jul 2026 17:13:49 +0800 Subject: [PATCH 09/13] fix lint and NixOS UT --- .github/workflows/citest.yaml | 79 +++ .github/workflows/publish.yaml | 22 + .pre-commit-config.yaml | 16 + pyproject.toml | 24 +- src/modelscope_hub/_download.py | 83 +-- src/modelscope_hub/_git.py | 6 +- src/modelscope_hub/_legacy_api.py | 17 +- src/modelscope_hub/_openapi.py | 56 +- src/modelscope_hub/_upload.py | 508 +++++++----------- src/modelscope_hub/agent/__init__.py | 1 + src/modelscope_hub/agent/_api.py | 217 +++++--- src/modelscope_hub/api.py | 18 +- src/modelscope_hub/cli/agent.py | 123 ++--- src/modelscope_hub/cli/base.py | 28 +- src/modelscope_hub/cli/cache.py | 11 +- src/modelscope_hub/cli/compat.py | 16 +- src/modelscope_hub/cli/deploy.py | 15 +- src/modelscope_hub/cli/download.py | 27 +- src/modelscope_hub/cli/login.py | 8 +- src/modelscope_hub/cli/main.py | 21 +- src/modelscope_hub/cli/mcp.py | 64 ++- src/modelscope_hub/cli/repo.py | 92 ++-- src/modelscope_hub/cli/secret.py | 18 +- src/modelscope_hub/cli/upload.py | 13 +- src/modelscope_hub/compat/__init__.py | 6 +- src/modelscope_hub/compat/file_download.py | 20 +- src/modelscope_hub/compat/hub_api.py | 119 ++-- .../compat/snapshot_download.py | 25 +- src/modelscope_hub/config.py | 52 +- src/modelscope_hub/constants.py | 276 +++++++--- src/modelscope_hub/errors.py | 93 ++-- src/modelscope_hub/types.py | 7 +- src/modelscope_hub/utils/file_utils.py | 6 +- src/modelscope_hub/utils/format.py | 5 +- src/modelscope_hub/utils/media.py | 4 +- src/modelscope_hub/utils/time_utils.py | 3 +- tests/agent/test_agent_cli.py | 79 ++- tests/agent/test_anonymous_download.py | 52 +- tests/cli/conftest.py | 55 +- tests/cli/run_all.py | 5 +- tests/cli/test_base.py | 5 +- tests/cli/test_cache.py | 1 - tests/cli/test_compat.py | 38 +- tests/cli/test_compat_cache_dir.py | 27 +- tests/cli/test_compat_revision.py | 44 +- tests/cli/test_deploy.py | 64 ++- tests/cli/test_download.py | 149 +++-- tests/cli/test_login.py | 25 +- tests/cli/test_main.py | 19 +- tests/cli/test_mcp.py | 140 ++++- tests/cli/test_openapi.py | 68 ++- tests/cli/test_repo.py | 465 +++++++++++----- tests/cli/test_secret.py | 29 +- tests/cli/test_upload.py | 123 +++-- tests/conftest.py | 9 +- tests/integration/conftest.py | 1 + tests/integration/run_all.py | 8 +- tests/integration/test_dataset_ops.py | 24 +- tests/integration/test_openapi.py | 1 + tests/integration/test_remote_file_ops.py | 5 +- tests/integration/test_remote_repo.py | 5 +- tests/integration/test_sdk_api.py | 107 +++- tests/test_compat_get_model_files.py | 26 +- tests/test_compat_snapshot_download.py | 9 +- tests/test_config_token.py | 1 + tests/test_legacy_cache_detection.py | 7 +- tests/test_upload_lfs_gate.py | 20 +- tests/test_utils.py | 13 +- 68 files changed, 2340 insertions(+), 1383 deletions(-) create mode 100644 .github/workflows/citest.yaml create mode 100644 .pre-commit-config.yaml diff --git a/.github/workflows/citest.yaml b/.github/workflows/citest.yaml new file mode 100644 index 0000000..a023025 --- /dev/null +++ b/.github/workflows/citest.yaml @@ -0,0 +1,79 @@ +name: citest + +on: + push: + branches: + - main + - "release/**" + paths-ignore: + - "docs/**" + - "temp/**" + - "README.md" + - "LICENSE" + - ".github/workflows/publish.yaml" + + pull_request: + paths-ignore: + - "docs/**" + - "temp/**" + - "README.md" + - "LICENSE" + - ".github/workflows/publish.yaml" + +concurrency: + group: ${{ github.workflow }}-${{ github.ref }} + cancel-in-progress: true + +jobs: + unittest: + # Mock mode: no credentials, no network calls to the Hub. This mirrors + # distro build sandboxes (NixOS/FreeBSD) so a green run here means the + # test suite also passes when downstream packagers run it. + runs-on: ubuntu-latest + timeout-minutes: 30 + strategy: + fail-fast: false + matrix: + # 3.10 = oldest supported, 3.14 = used by NixOS packaging. + python-version: ["3.10", "3.12", "3.14"] + env: + MODELSCOPE_RUN_REMOTE_TESTS: "false" + steps: + - name: Checkout + uses: actions/checkout@v4 + + - name: Set up Python ${{ matrix.python-version }} + uses: actions/setup-python@v5 + with: + python-version: ${{ matrix.python-version }} + + - name: Install package with dev dependencies + run: pip install -e ".[dev]" + + - name: Run unit tests (mock mode) + run: pytest tests/ --ignore=tests/integration -q + + lint: + # Hard gate: the ruff/mypy backlog was cleared, keep it at zero. + runs-on: ubuntu-latest + timeout-minutes: 10 + steps: + - name: Checkout + uses: actions/checkout@v4 + + - name: Set up Python + uses: actions/setup-python@v5 + with: + python-version: "3.13" + + - name: Install package with dev dependencies + run: pip install -e ".[dev]" + + - name: Ruff check + run: ruff check src/ tests/ + + - name: Ruff format check + run: ruff format --check src/ tests/ + + - name: Mypy + run: mypy src/modelscope_hub/ diff --git a/.github/workflows/publish.yaml b/.github/workflows/publish.yaml index 6caa1a5..4af3748 100644 --- a/.github/workflows/publish.yaml +++ b/.github/workflows/publish.yaml @@ -10,7 +10,29 @@ concurrency: cancel-in-progress: true jobs: + test-gate: + # Release gate: the mock-mode suite must pass before anything is built + # or uploaded. Same environment shape as downstream distro sandboxes. + runs-on: ubuntu-latest + timeout-minutes: 30 + env: + MODELSCOPE_RUN_REMOTE_TESTS: "false" + steps: + - uses: actions/checkout@v4 + + - name: Set up Python 3.13 + uses: actions/setup-python@v5 + with: + python-version: '3.13' + + - name: Install package with dev dependencies + run: pip install -e ".[dev]" + + - name: Run unit tests (mock mode) + run: pytest tests/ --ignore=tests/integration -q + build-n-publish: + needs: test-gate runs-on: ubuntu-22.04 steps: - uses: actions/checkout@v4 diff --git a/.pre-commit-config.yaml b/.pre-commit-config.yaml new file mode 100644 index 0000000..11cb344 --- /dev/null +++ b/.pre-commit-config.yaml @@ -0,0 +1,16 @@ +# Local pre-commit hooks. Optional but recommended: CI (citest.yaml) enforces +# the same checks as a hard gate, this just gives faster feedback. +# +# Setup: +# pip install pre-commit && pre-commit install +repos: + - repo: https://github.com/astral-sh/ruff-pre-commit + rev: v0.15.11 + hooks: + - id: ruff-check + args: [--fix] + files: ^(src|tests)/ + - id: ruff-format + files: ^(src|tests)/ +# mypy is intentionally left to CI: a full run over src/ takes too long for a +# commit hook and needs the package's own dependencies installed. diff --git a/pyproject.toml b/pyproject.toml index 898b356..eda3b10 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -52,12 +52,24 @@ modelscope_hub = ["py.typed"] version = {attr = "modelscope_hub.version.__version__"} [tool.ruff] -target-version = "py313" +# Must match requires-python (>=3.10): with py313 the UP rules would +# auto-rewrite code into 3.12+ syntax and silently break 3.10 support. +target-version = "py310" line-length = 120 [tool.ruff.lint] select = ["E", "F", "I", "N", "W", "UP"] +[tool.ruff.lint.per-file-ignores] +# Public exception names mirror the legacy SDK; renaming breaks downstream compat. +"src/modelscope_hub/errors.py" = ["N818"] +# Function-scope constants deliberately use CONST_CASE. +"src/modelscope_hub/api.py" = ["N806"] +"src/modelscope_hub/types.py" = ["N806"] +"src/modelscope_hub/compat/hub_api.py" = ["N806"] +# unittest.mock patch decorators inject mock classes; CamelCase args mirror the patched target. +"tests/**" = ["N803"] + [tool.pytest.ini_options] testpaths = ["tests"] markers = [ @@ -66,6 +78,12 @@ markers = [ addopts = "-v --tb=short" [tool.mypy] -python_version = "3.13" +python_version = "3.10" strict = false -warn_return_any = true +# The HTTP boundary (requests -> JSON) is dynamically typed; until response +# models are introduced, warning on Any-returns only produces noise. +warn_return_any = false + +[[tool.mypy.overrides]] +module = ["tqdm.*", "modelscope.*"] +ignore_missing_imports = true diff --git a/src/modelscope_hub/_download.py b/src/modelscope_hub/_download.py index 4904f83..d1fa1b9 100644 --- a/src/modelscope_hub/_download.py +++ b/src/modelscope_hub/_download.py @@ -24,10 +24,9 @@ import hashlib import io import os -import time import re import threading - +import time import uuid from concurrent.futures import ThreadPoolExecutor, as_completed from pathlib import Path @@ -46,9 +45,9 @@ DOWNLOAD_RETRY_TIMES, DOWNLOAD_TIMEOUT, ENV_FILE_LOCK, + ENV_INTER_CLOUD_REGIONS, ENV_INTRA_CLOUD_ACCELERATION, ENV_INTRA_CLOUD_REGION, - ENV_INTER_CLOUD_REGIONS, ) from .errors import ( CacheNotFound, @@ -62,8 +61,8 @@ from .utils.logger import get_logger if TYPE_CHECKING: - from .config import HubConfig from ._legacy_api import LegacyClient + from .config import HubConfig logger = get_logger("download") @@ -148,7 +147,8 @@ def _optional_file_lock(lock_path: Path | None, *, enabled: bool = True): if age >= _STALE_LOCK_SECONDS: logger.warning( "Removing possibly stale SoftFileLock (age=%.0fs): %s", - age, lock_path, + age, + lock_path, ) try: lock_path.unlink(missing_ok=True) @@ -174,7 +174,8 @@ def _optional_file_lock(lock_path: Path | None, *, enabled: bool = True): if exc.errno in (errno.ESTALE, errno.ENOENT, getattr(errno, "EREMOTEIO", -1)): logger.warning( "OSError (errno=%d) on %s, falling back to SoftFileLock.", - exc.errno, lock_path, + exc.errno, + lock_path, ) lock = SoftFileLock(str(lock_path), timeout=default_interval) is_soft = True @@ -198,6 +199,7 @@ def _optional_file_lock(lock_path: Path | None, *, enabled: bool = True): def _file_lock_enabled() -> bool: """Check whether file locking is enabled via environment.""" from .constants import _env_bool + return _env_bool(ENV_FILE_LOCK, True) @@ -250,8 +252,11 @@ def _download_part_with_retry(params: tuple) -> None: get_headers["Range"] = f"bytes={download_start}-{end}" with open(part_file_name, "ab+") as f: r = requests.get( - url, stream=True, headers=get_headers, - cookies=cookies, timeout=DOWNLOAD_TIMEOUT, + url, + stream=True, + headers=get_headers, + cookies=cookies, + timeout=DOWNLOAD_TIMEOUT, ) r.raise_for_status() for chunk in r.iter_content(chunk_size=DOWNLOAD_CHUNK_SIZE): @@ -326,7 +331,7 @@ class DownloadManager: Dependencies are injected via constructor to keep this class testable. """ - def __init__(self, legacy_client: "LegacyClient", config: "HubConfig") -> None: + def __init__(self, legacy_client: LegacyClient, config: HubConfig) -> None: self._client = legacy_client self._config = config self._cached_region: str | None = None @@ -380,9 +385,7 @@ def _get(url: str, timeout: float): # OSS internal endpoint hostname pattern: # .oss-internal.aliyuncs.com # e.g. modelhub-cn-hangzhou.oss-cn-hangzhou-internal.aliyuncs.com - _OSS_INTERNAL_RE = re.compile( - r".*\.oss.*-internal\.aliyuncs\.com$" - ) + _OSS_INTERNAL_RE = re.compile(r".*\.oss.*-internal\.aliyuncs\.com$") @staticmethod def _is_oss_internal_url(url: str) -> bool: @@ -407,8 +410,11 @@ def _probe_redirect_url( """Send a HEAD request without following redirects to get the 302 Location.""" try: r = requests.head( - url, headers=headers, cookies=cookies, - allow_redirects=False, timeout=timeout, + url, + headers=headers, + cookies=cookies, + allow_redirects=False, + timeout=timeout, ) if r.status_code in (301, 302, 303, 307, 308): return r.headers.get("Location", "") @@ -419,6 +425,7 @@ def _probe_redirect_url( def _get_inter_cloud_regions(self) -> list[str]: """Read the inter-cloud peer region list from the environment.""" from .constants import _env + raw = _env(ENV_INTER_CLOUD_REGIONS, "INTER_CLOUD_ACCELERATION_REGIONS") or "" return [r.strip().lower() for r in raw.split(",") if r.strip()] @@ -586,7 +593,11 @@ def download_file( for attempt in range(DOWNLOAD_HASH_RETRY_TIMES): self._download_with_resume( - repo_id, repo_type, file_path, revision, target, + repo_id, + repo_type, + file_path, + revision, + target, file_size=file_size, user_agent=user_agent, progress_callbacks=progress_callbacks, @@ -602,7 +613,9 @@ def download_file( if attempt < DOWNLOAD_HASH_RETRY_TIMES - 1: logger.warning( "Hash validation failed for %s, retrying (%d/%d)", - file_path, attempt + 1, DOWNLOAD_HASH_RETRY_TIMES, + file_path, + attempt + 1, + DOWNLOAD_HASH_RETRY_TIMES, ) target.unlink(missing_ok=True) else: @@ -663,7 +676,8 @@ def download_repo( 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, + "Found legacy cache at %s, reusing.", + legacy, ) output_dir = legacy local_dir = legacy @@ -673,9 +687,7 @@ def download_repo( if local_files_only: if any(output_dir.iterdir()): - logger.warning( - "Cannot confirm the cached file is for revision: %s", revision - ) + 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" @@ -913,11 +925,7 @@ def _download_with_resume( progress_callbacks: list[type[ProgressCallback]] | None = None, ) -> Path: """Download a file with HTTP Range resume support and retry.""" - use_parallel = ( - file_size is not None - and file_size > DOWNLOAD_PARALLEL_THRESHOLD - and DOWNLOAD_PARALLELS > 1 - ) + use_parallel = file_size is not None and file_size > DOWNLOAD_PARALLEL_THRESHOLD and DOWNLOAD_PARALLELS > 1 download_headers = self._build_download_headers(user_agent) @@ -936,20 +944,26 @@ def _download_with_resume( # other threads wait instead of issuing redundant HEAD reqs. try: probe_url = self._client.get_download_url( - repo_id, repo_type, file_path, revision, + repo_id, + repo_type, + file_path, + revision, ) cookies = None if self._client.token: cookies = {"m_session_id": self._client.token} download_headers, source = self._resolve_inter_region_headers( - probe_url, download_headers, cookies, + probe_url, + download_headers, + cookies, peer_regions=peer_regions, ) resolved_region = download_headers.get("x-aliyun-region-id") self._inter_region_cache[cache_key] = (resolved_region, source) except Exception as exc: logger.warning( - "Failed to resolve inter-region acceleration: %s. Falling back to default.", exc, + "Failed to resolve inter-region acceleration: %s. Falling back to default.", + exc, ) self._inter_region_cache[cache_key] = (None, "default") cached = (None, "default") @@ -962,12 +976,18 @@ def _download_with_resume( if resolved_region is not None: download_headers["x-aliyun-region-id"] = resolved_region source_prefix = { - "local": "\u26a1 ", "peer": "\u21c4 ", "default": " ", + "local": "\u26a1 ", + "peer": "\u21c4 ", + "default": " ", }[source] if use_parallel: + assert file_size is not None # guaranteed by use_parallel above url = self._client.get_download_url( - repo_id, repo_type, file_path, revision, + repo_id, + repo_type, + file_path, + revision, ) cookies = None if self._client.token: @@ -1075,7 +1095,6 @@ def verify_file(self, file_path: Path, expected_sha256: str) -> bool: actual = compute_hash(file_path, "sha256") if actual != expected_sha256: raise FileIntegrityError( - f"Hash mismatch for {file_path.name}: " - f"expected {expected_sha256[:16]}..., got {actual[:16]}..." + f"Hash mismatch for {file_path.name}: expected {expected_sha256[:16]}..., got {actual[:16]}..." ) return True diff --git a/src/modelscope_hub/_git.py b/src/modelscope_hub/_git.py index 74b2810..c1b2d76 100644 --- a/src/modelscope_hub/_git.py +++ b/src/modelscope_hub/_git.py @@ -73,8 +73,7 @@ def _run(cls, *args: str, cwd: Path | str | None = None) -> subprocess.Completed result = subprocess.run( cmd, cwd=cwd, - stdout=subprocess.PIPE, - stderr=subprocess.PIPE, + capture_output=True, text=True, env=env, ) @@ -132,8 +131,7 @@ def clone( # Clone may succeed but hook fails — check if .git exists if (target_dir / ".git").is_dir(): logger.warning( - "Clone exited non-zero but repository exists at %s. " - "Likely caused by a post-clone hook.", + "Clone exited non-zero but repository exists at %s. Likely caused by a post-clone hook.", target_dir, ) else: diff --git a/src/modelscope_hub/_legacy_api.py b/src/modelscope_hub/_legacy_api.py index 5fa7ad1..b460e33 100644 --- a/src/modelscope_hub/_legacy_api.py +++ b/src/modelscope_hub/_legacy_api.py @@ -15,7 +15,7 @@ from __future__ import annotations import uuid -from typing import Any, BinaryIO, IO, Union +from typing import IO, Any, BinaryIO from urllib.parse import quote_plus, urlparse import requests @@ -26,10 +26,10 @@ API_MAX_RETRIES, API_TIMEOUT, LEGACY_API_PREFIX, - RepoType, UPLOAD_BLOB_CONNECT_TIMEOUT, UPLOAD_BLOB_READ_TIMEOUT, UPLOAD_RETRY_ALLOWED_METHODS, + RepoType, ) from .errors import InvalidParameter, NetworkError, RequestTimeoutError, ServerError, raise_for_status from .utils.logger import get_logger @@ -208,7 +208,7 @@ def _json_data(self, resp: requests.Response) -> Any: # ------------------------------------------------------------------ # Auth # ------------------------------------------------------------------ - def login(self, access_token: str) -> tuple[dict, "requests.cookies.RequestsCookieJar"]: + def login(self, access_token: str) -> tuple[dict, requests.cookies.RequestsCookieJar]: """Authenticate via access token and return (user_data, cookies). POST /api/v1/login @@ -332,7 +332,7 @@ def list_repo_files( return data # Sometimes wrapped: {"Data": {"Files": [...]}} if isinstance(data, dict): - return data.get("Files", data.get("files", [])) + return data.get("Files") or data.get("files") or [] return [] def list_dataset_files_paginated( @@ -368,7 +368,7 @@ def list_dataset_files_paginated( if isinstance(data, list): files = data elif isinstance(data, dict): - files = data.get("Files", data.get("files", [])) + files = data.get("Files") or data.get("files") or [] else: files = [] @@ -415,7 +415,9 @@ def list_revisions_detail( if end_time is not None: params["EndTime"] = end_time resp = self._request( - "GET", f"{segment}/{repo_id}/revisions", params=params or None, + "GET", + f"{segment}/{repo_id}/revisions", + params=params or None, ) data = self._json_data(resp) if isinstance(data, dict): @@ -547,7 +549,7 @@ def validate_blobs( def upload_blob( self, upload_url: str, - data: Union[str, bytes, BinaryIO, IO[bytes]], + data: str | bytes | BinaryIO | IO[bytes] | Any, size: int, *, headers: dict[str, str] | None = None, @@ -594,6 +596,7 @@ def upload_blob( return {} if isinstance(body, dict) and body.get("Code") not in (200, "200", None): from .errors import APIError + raise APIError( body.get("Message") or body.get("message") or f"Blob upload failed (Code={body.get('Code')})", status_code=resp.status_code, diff --git a/src/modelscope_hub/_openapi.py b/src/modelscope_hub/_openapi.py index 690f317..0306560 100644 --- a/src/modelscope_hub/_openapi.py +++ b/src/modelscope_hub/_openapi.py @@ -19,8 +19,9 @@ import random import time +from collections.abc import Iterable, Mapping from pathlib import Path -from typing import Any, BinaryIO, Iterable, Mapping +from typing import Any, BinaryIO from urllib.parse import urljoin, urlsplit import requests @@ -57,7 +58,9 @@ # Errors that warrant a transparent retry. _RETRYABLE_EXC: tuple[type[BaseException], ...] = ( - NetworkError, ServerError, RateLimitError, + NetworkError, + ServerError, + RateLimitError, ) JSON = dict[str, Any] @@ -94,8 +97,7 @@ def __init__( self._config = config or get_default_config() self._session = session or requests.Session() self._timeout: float | tuple[float, float] = ( - float(timeout) if timeout is not None - else (float(API_CONNECT_TIMEOUT), float(API_TIMEOUT)) + float(timeout) if timeout is not None else (float(API_CONNECT_TIMEOUT), float(API_TIMEOUT)) ) self._max_retries = int(max_retries) if max_retries is not None else int(API_MAX_RETRIES) @@ -106,7 +108,7 @@ def close(self) -> None: """Release the underlying HTTP session.""" self._session.close() - def __enter__(self) -> "OpenAPIClient": + def __enter__(self) -> OpenAPIClient: return self def __exit__(self, *_exc: object) -> None: @@ -139,10 +141,16 @@ def request( When *unwrap* is ``False`` the raw :class:`requests.Response` is returned. """ return self._request( - method, path, + method, + path, url=url, - params=params, json_body=json_body, data=data, files=files, - headers=headers, require_token=require_token, unwrap=unwrap, + params=params, + json_body=json_body, + data=data, + files=files, + headers=headers, + require_token=require_token, + unwrap=unwrap, timeout=timeout, ) @@ -152,7 +160,7 @@ def request( @property def base_url(self) -> str: """Fully-qualified OpenAPI base URL, including trailing slash.""" - return f"{self._config.endpoint.rstrip('/')}{OPENAPI_PREFIX}/" + return f"{(self._config.endpoint or '').rstrip('/')}{OPENAPI_PREFIX}/" def _url(self, path: str) -> str: # ``urljoin`` treats absolute leading slashes as roots, which would @@ -172,9 +180,7 @@ def _auth_headers(self, *, require_token: bool = False) -> dict[str, str]: token = self._resolve_token() if not token: if require_token: - raise AuthenticationError( - "Missing API token. Call HubApi.login(...) or set MODELSCOPE_API_TOKEN." - ) + raise AuthenticationError("Missing API token. Call HubApi.login(...) or set MODELSCOPE_API_TOKEN.") return {} return {"Authorization": f"Bearer {token}"} @@ -346,7 +352,11 @@ def _request( backoff = min(2 ** (attempt - 1), 16) + random.uniform(0, 0.5) _logger.debug( "Retrying %s %s after %s (attempt %d/%d)", - method_upper, final_url, last_exc, attempt, attempts, + method_upper, + final_url, + last_exc, + attempt, + attempts, ) time.sleep(backoff) @@ -397,9 +407,7 @@ def list_models( ``custom_tag``, ``license``, ``deploy``. """ if page_number * page_size > 3000: - raise InvalidParameter( - f"page_number * page_size must be <= 3000 (got {page_number * page_size})." - ) + raise InvalidParameter(f"page_number * page_size must be <= 3000 (got {page_number * page_size}).") params = self._merge_params( { "search": search, @@ -431,9 +439,7 @@ def list_datasets( ) -> JSON: """``GET /datasets`` — list datasets. Filter keys: ``task``, ``license``.""" if page_number * page_size > 3000: - raise InvalidParameter( - f"page_number * page_size must be <= 3000 (got {page_number * page_size})." - ) + raise InvalidParameter(f"page_number * page_size must be <= 3000 (got {page_number * page_size}).") params = self._merge_params( { "search": search, @@ -518,9 +524,7 @@ def list_skills( ``owner``. """ if page_number * page_size > 3000: - raise InvalidParameter( - f"page_number * page_size must be <= 3000 (got {page_number * page_size})." - ) + raise InvalidParameter(f"page_number * page_size must be <= 3000 (got {page_number * page_size}).") params = self._merge_params( { "search": search, @@ -696,9 +700,7 @@ def get_mcp_server( ) -> JSON: """``GET /mcp/servers/{id}`` — fetch a single MCP server's manifest.""" params = self._merge_params({"get_operational_url": get_operational_url}) - return self._request( - "GET", f"/mcp/servers/{server_id}", params=params, require_token=False - ) + return self._request("GET", f"/mcp/servers/{server_id}", params=params, require_token=False) def deploy_mcp_server( self, @@ -706,7 +708,9 @@ def deploy_mcp_server( payload: DeployMcpServerPayload | Mapping[str, Any] | None = None, ) -> JSON: """``POST /mcp/servers/{id}/deploy`` — deploy an MCP server for the caller.""" - body = dict(payload or {}) + # Drop explicit None values so they never reach the wire, then apply + # the platform default transport. + body = {k: v for k, v in dict(payload or {}).items() if v is not None} body.setdefault("transport_type", "sse") return self._request( "POST", diff --git a/src/modelscope_hub/_upload.py b/src/modelscope_hub/_upload.py index 5905814..28bbd66 100644 --- a/src/modelscope_hub/_upload.py +++ b/src/modelscope_hub/_upload.py @@ -20,13 +20,12 @@ import io import json import os -import re import tempfile import threading import time from concurrent.futures import ThreadPoolExecutor, as_completed from pathlib import Path -from typing import TYPE_CHECKING, Any, BinaryIO, IO, Union +from typing import IO, TYPE_CHECKING, Any, BinaryIO from tqdm.auto import tqdm @@ -78,7 +77,7 @@ logger = get_logger("upload") -PathOrFileObj = Union[str, Path, bytes, BinaryIO, IO[bytes]] +PathOrFileObj = str | Path | bytes | BinaryIO | IO[bytes] _TRACKER_VERSION = 3 @@ -91,8 +90,7 @@ class _CountedReadStream: """File wrapper that counts bytes read and updates a progress bar.""" - def __init__(self, file_obj: Any, expected_size: int, - pbar: Any, chunk_size: int) -> None: + def __init__(self, file_obj: Any, expected_size: int, pbar: Any, chunk_size: int) -> None: self._file = file_obj self._expected_size = expected_size self._pbar = pbar @@ -221,19 +219,19 @@ def is_retryable(cls, category: str) -> bool: _CATEGORY_BY_ERROR_CODE: dict[str, str] = { - "E1001": _ErrorCategory.TRANSIENT_NETWORK, # timeout - "E1002": _ErrorCategory.TRANSIENT_SERVER, # server error - "E1003": _ErrorCategory.TRANSIENT_SERVER, # storage error - "E1020": _ErrorCategory.TRANSIENT_NETWORK, # network/connection error - "E1021": _ErrorCategory.THROTTLED, # rate limit - "E1022": _ErrorCategory.FILE_INVALID, # cache error - "E2020": _ErrorCategory.TRANSIENT_SERVER, # file integrity (auto-retry) - "E3001": _ErrorCategory.AUTH_FAILED, # authentication - "E3002": _ErrorCategory.AUTH_FAILED, # permission - "E3020": _ErrorCategory.NOT_FOUND, # not exist - "E3021": _ErrorCategory.FILE_INVALID, # invalid parameter - "E3023": _ErrorCategory.FILE_INVALID, # not supported - "E9001": _ErrorCategory.UNKNOWN, # unknown/fallback + "E1001": _ErrorCategory.TRANSIENT_NETWORK, # timeout + "E1002": _ErrorCategory.TRANSIENT_SERVER, # server error + "E1003": _ErrorCategory.TRANSIENT_SERVER, # storage error + "E1020": _ErrorCategory.TRANSIENT_NETWORK, # network/connection error + "E1021": _ErrorCategory.THROTTLED, # rate limit + "E1022": _ErrorCategory.FILE_INVALID, # cache error + "E2020": _ErrorCategory.TRANSIENT_SERVER, # file integrity (auto-retry) + "E3001": _ErrorCategory.AUTH_FAILED, # authentication + "E3002": _ErrorCategory.AUTH_FAILED, # permission + "E3020": _ErrorCategory.NOT_FOUND, # not exist + "E3021": _ErrorCategory.FILE_INVALID, # invalid parameter + "E3023": _ErrorCategory.FILE_INVALID, # not supported + "E9001": _ErrorCategory.UNKNOWN, # unknown/fallback } @@ -305,8 +303,7 @@ def get_hash(self, rel_path: str, mtime: float, size: int) -> dict | None: "file_size": entry["size"], } - def put_hash(self, rel_path: str, mtime: float, size: int, - hash_info: dict) -> None: + def put_hash(self, rel_path: str, mtime: float, size: int, hash_info: dict) -> None: key = self._make_key(rel_path, mtime, size) with self._lock: entry = self._files.get(key, {}) @@ -334,9 +331,7 @@ def mark_uploaded(self, rel_path: str, mtime: float, size: int) -> None: self._files[key]["status"] = FileStatus.UPLOADED self._dirty = True - def mark_committed_batch( - self, file_keys: list[tuple[str, float, int]] - ) -> None: + def mark_committed_batch(self, file_keys: list[tuple[str, float, int]]) -> None: with self._lock: for rel_path, mtime, size in file_keys: key = self._make_key(rel_path, mtime, size) @@ -344,8 +339,7 @@ def mark_committed_batch( self._files[key]["status"] = FileStatus.COMMITTED self._dirty = True - def mark_failed(self, rel_path: str, mtime: float, size: int, - error_type: str = "") -> None: + def mark_failed(self, rel_path: str, mtime: float, size: int, error_type: str = "") -> None: key = self._make_key(rel_path, mtime, size) with self._lock: if key in self._files: @@ -371,9 +365,7 @@ def save(self) -> None: self._dirty = False try: self._path.parent.mkdir(parents=True, exist_ok=True) - fd, tmp_path = tempfile.mkstemp( - dir=str(self._path.parent), suffix=".tmp" - ) + fd, tmp_path = tempfile.mkstemp(dir=str(self._path.parent), suffix=".tmp") try: with os.fdopen(fd, "w", encoding="utf-8") as f: json.dump(data, f, ensure_ascii=False) @@ -398,7 +390,7 @@ def _load(self) -> None: self._check_legacy_progress() return try: - with open(self._path, "r") as f: + with open(self._path) as f: data = json.load(f) except (json.JSONDecodeError, OSError) as e: logger.warning("Failed to load upload tracker, starting fresh: %s", e) @@ -412,21 +404,19 @@ def _load(self) -> None: stored_repo = data.get("repo_id", "") if stored_repo and stored_repo != self._repo_id: logger.warning( - "Tracker repo_id mismatch (cached: %s, current: %s), " - "ignoring stale tracker.", - stored_repo, self._repo_id, + "Tracker repo_id mismatch (cached: %s, current: %s), ignoring stale tracker.", + stored_repo, + self._repo_id, ) return self._files = data.get("files", {}) - committed_count = sum( - 1 for e in self._files.values() - if e.get("status") == FileStatus.COMMITTED - ) + committed_count = sum(1 for e in self._files.values() if e.get("status") == FileStatus.COMMITTED) if committed_count > 0: logger.info( "Upload tracker loaded: %d entries, %d committed.", - len(self._files), committed_count, + len(self._files), + committed_count, ) self._check_legacy_progress() @@ -450,8 +440,7 @@ def _check_legacy_progress(self) -> None: legacy_path = self._path.parent / UPLOAD_LEGACY_PROGRESS_FILE if legacy_path.exists(): logger.warning( - "Legacy upload progress file detected: %s. " - "This file is no longer used. You may delete it safely.", + "Legacy upload progress file detected: %s. This file is no longer used. You may delete it safely.", legacy_path, ) @@ -462,8 +451,7 @@ class NullTracker: def get_hash(self, rel_path: str, mtime: float, size: int) -> None: return None - def put_hash(self, rel_path: str, mtime: float, size: int, - hash_info: dict) -> None: + def put_hash(self, rel_path: str, mtime: float, size: int, hash_info: dict) -> None: pass def is_committed(self, rel_path: str, mtime: float, size: int) -> bool: @@ -478,8 +466,7 @@ def mark_uploaded(self, rel_path: str, mtime: float, size: int) -> None: def mark_committed_batch(self, file_keys: list) -> None: pass - def mark_failed(self, rel_path: str, mtime: float, size: int, - error_type: str = "") -> None: + def mark_failed(self, rel_path: str, mtime: float, size: int, error_type: str = "") -> None: pass def save(self) -> None: @@ -499,23 +486,15 @@ class BatchTracker: def __init__(self, total_files: int, batch_size: int) -> None: self._batch_size = batch_size - self._num_batches = ( - (total_files - 1) // batch_size + 1 if total_files > 0 else 0 - ) - self._batch_results: list[list[dict]] = [ - [] for _ in range(self._num_batches) - ] - self._batch_failures: list[list[tuple]] = [ - [] for _ in range(self._num_batches) - ] + self._num_batches = (total_files - 1) // batch_size + 1 if total_files > 0 else 0 + self._batch_results: list[list[dict]] = [[] for _ in range(self._num_batches)] + self._batch_failures: list[list[tuple]] = [[] for _ in range(self._num_batches)] self._batch_expected: list[int] = [] for i in range(self._num_batches): start = i * batch_size end = min(start + batch_size, total_files) self._batch_expected.append(end - start) - self._batch_events: list[threading.Event] = [ - threading.Event() for _ in range(self._num_batches) - ] + self._batch_events: list[threading.Event] = [threading.Event() for _ in range(self._num_batches)] self._lock = threading.Lock() @property @@ -532,8 +511,7 @@ def record_success(self, file_index: int, result: dict) -> None: if self._is_batch_complete(idx): self._batch_events[idx].set() - def record_failure(self, file_index: int, item: tuple, - error: Exception) -> None: + def record_failure(self, file_index: int, item: tuple, error: Exception) -> None: idx = self.batch_index(file_index) with self._lock: self._batch_failures[idx].append((item, error)) @@ -547,9 +525,7 @@ def mark_file_skipped(self, file_index: int) -> None: if self._is_batch_complete(idx): self._batch_events[idx].set() - def wait_for_batch( - self, batch_idx: int - ) -> tuple[list[dict], list[tuple]]: + def wait_for_batch(self, batch_idx: int) -> tuple[list[dict], list[tuple]]: self._batch_events[batch_idx].wait() with self._lock: return ( @@ -558,10 +534,7 @@ def wait_for_batch( ) def _is_batch_complete(self, batch_idx: int) -> bool: - count = ( - len(self._batch_results[batch_idx]) - + len(self._batch_failures[batch_idx]) - ) + count = len(self._batch_results[batch_idx]) + len(self._batch_failures[batch_idx]) return count >= self._batch_expected[batch_idx] @@ -575,9 +548,9 @@ class UploadManager: def __init__( self, - legacy_client: "LegacyClient", - config: "HubConfig", - openapi_client: "OpenAPIClient | None" = None, + legacy_client: LegacyClient, + config: HubConfig, + openapi_client: OpenAPIClient | None = None, *, create_repo_fn: Any = None, ) -> None: @@ -607,9 +580,7 @@ def upload_file( raise InvalidParameter("Path or file object cannot be None!") if isinstance(path_or_fileobj, (str, Path)): - path_or_fileobj = os.path.abspath( - os.path.expanduser(str(path_or_fileobj)) - ) + path_or_fileobj = os.path.abspath(os.path.expanduser(str(path_or_fileobj))) path_in_repo = path_in_repo or os.path.basename(path_or_fileobj) else: if not path_in_repo: @@ -706,15 +677,9 @@ def upload_folder( ignore_patterns += DEFAULT_IGNORE_PATTERNS if allow_patterns is not None: - ignore_patterns = [ - p for p in ignore_patterns if p not in allow_patterns - ] - - commit_message = ( - commit_message - if commit_message is not None - else f"Upload to {repo_id} on ModelScope hub" - ) + ignore_patterns = [p for p in ignore_patterns if p not in allow_patterns] + + commit_message = commit_message if commit_message is not None else f"Upload to {repo_id} on ModelScope hub" commit_description = commit_description or "Uploading files" # Exclude internal cache files from upload @@ -762,22 +727,17 @@ def upload_folder( commit_batch_size = _calculate_adaptive_batch_size(len(sorted_files)) logger.info( "Adaptive batch size: %d (for %d files)", - commit_batch_size, len(sorted_files), + commit_batch_size, + len(sorted_files), ) else: - commit_batch_size = ( - UPLOAD_COMMIT_BATCH_SIZE - if UPLOAD_COMMIT_BATCH_SIZE > 0 - else len(sorted_files) - ) + commit_batch_size = UPLOAD_COMMIT_BATCH_SIZE if UPLOAD_COMMIT_BATCH_SIZE > 0 else len(sorted_files) # Initialize tracker folder_path_resolved = Path(folder_path).resolve() if use_cache: cache_path = folder_path_resolved / UPLOAD_CACHE_FILE - tracker: UploadTracker | NullTracker = UploadTracker( - cache_path, repo_id=repo_id - ) + tracker: UploadTracker | NullTracker = UploadTracker(cache_path, repo_id=repo_id) else: tracker = NullTracker() batch_tracker = BatchTracker(len(sorted_files), commit_batch_size) @@ -788,20 +748,17 @@ def upload_folder( for file_idx, (file_path_in_repo, file_path) in enumerate(sorted_files): try: st = os.stat(file_path) - if tracker.is_committed( - file_path_in_repo, st.st_mtime, st.st_size - ): + if tracker.is_committed(file_path_in_repo, st.st_mtime, st.st_size): skipped_indices.add(file_idx) batch_tracker.mark_file_skipped(file_idx) continue except OSError as e: logger.warning( "Cannot stat file %s, will re-upload: %s", - file_path_in_repo, e, + file_path_in_repo, + e, ) - files_to_upload.append( - (file_idx, (file_path_in_repo, file_path)) - ) + files_to_upload.append((file_idx, (file_path_in_repo, file_path))) # Batch pre-validation for LFS files with cached hashes pre_validated_map: dict[str, str | None] = {} @@ -810,14 +767,15 @@ def upload_folder( for file_idx, (file_path_in_repo, file_path) in files_to_upload: try: st = os.stat(file_path) - cached = tracker.get_hash( - file_path_in_repo, st.st_mtime, st.st_size - ) + cached = tracker.get_hash(file_path_in_repo, st.st_mtime, st.st_size) if cached is not None: if ( _upload_mode( - file_path_in_repo, cached["file_size"], repo_type, - ) == "lfs" + file_path_in_repo, + cached["file_size"], + repo_type, + ) + == "lfs" ): lfs_hash_info_map[file_idx] = (cached, st) continue @@ -825,19 +783,15 @@ def upload_folder( pass if lfs_hash_info_map: - objects = [ - {"oid": info["file_hash"], "size": info["file_size"]} - for info, _ in lfs_hash_info_map.values() - ] - validated = self._validate_blobs_batch( - repo_id=repo_id, repo_type=repo_type, objects=objects - ) + objects = [{"oid": info["file_hash"], "size": info["file_size"]} for info, _ in lfs_hash_info_map.values()] + validated = self._validate_blobs_batch(repo_id=repo_id, repo_type=repo_type, objects=objects) pre_validated_map = validated reused = sum(1 for v in validated.values() if v is None) logger.info( - "Pre-validated %d cached LFS hash(es): %d globally existing, " - "%d need upload.", - len(objects), reused, len(objects) - reused, + "Pre-validated %d cached LFS hash(es): %d globally existing, %d need upload.", + len(objects), + reused, + len(objects) - reused, ) skipped_count = len(skipped_indices) @@ -846,18 +800,20 @@ def upload_folder( logger.info( "Scan complete: %d total, %d committed (skip), %d to process.", - len(sorted_files), skipped_count, len(files_to_upload), + len(sorted_files), + skipped_count, + len(files_to_upload), ) logger.info( "Uploading %d file(s) in %d batch(es) of size %d (pipeline mode).", - len(files_to_upload), batch_tracker.num_batches, commit_batch_size, + len(files_to_upload), + batch_tracker.num_batches, + commit_batch_size, ) # Pipeline: upload workers - def _upload_worker( - file_idx: int, file_info: tuple, pre_validated: Any = None - ) -> None: + def _upload_worker(file_idx: int, file_info: tuple, pre_validated: Any = None) -> None: path_in_repo_w, file_path_w = file_info try: logger.debug("Uploading: %s ...", path_in_repo_w) @@ -885,7 +841,7 @@ def _upload_worker( try: with ThreadPoolExecutor(max_workers=max_workers) as executor: for file_idx, file_info in files_to_upload: - pv = None + pv: str | bool | None = None if file_idx in lfs_hash_info_map: cached_hash = lfs_hash_info_map[file_idx][0]["file_hash"] pv = pre_validated_map.get(cached_hash) @@ -901,16 +857,12 @@ def _upload_worker( disable=disable_tqdm, ): batch_start = batch_idx * commit_batch_size - batch_end = min( - batch_start + commit_batch_size, len(sorted_files) - ) - if all( - i in skipped_indices - for i in range(batch_start, batch_end) - ): + batch_end = min(batch_start + commit_batch_size, len(sorted_files)) + if all(i in skipped_indices for i in range(batch_start, batch_end)): logger.info( "Batch %d/%d fully committed, skipping.", - batch_idx + 1, num_batches, + batch_idx + 1, + num_batches, ) continue @@ -923,19 +875,16 @@ def _upload_worker( self._track_uploaded_batch(tracker, results) - operations = self._build_batch_operations( - results, repo_type - ) + operations = self._build_batch_operations(results, repo_type) if not operations: logger.error( "Batch %d/%d: all files failed, skipping commit.", - batch_idx + 1, num_batches, + batch_idx + 1, + num_batches, ) continue - batch_commit_message = ( - f"{commit_message} (batch {batch_idx + 1}/{num_batches})" - ) + batch_commit_message = f"{commit_message} (batch {batch_idx + 1}/{num_batches})" try: commit_info = self._commit_with_retry( repo_id=repo_id, @@ -948,14 +897,18 @@ def _upload_worker( all_results.extend(results) logger.info( "Batch %d/%d: committed %d file(s).", - batch_idx + 1, num_batches, len(results), + batch_idx + 1, + num_batches, + len(results), ) self._track_committed_batch(tracker, results) consecutive_failures = 0 except Exception as e: logger.error( "Batch %d/%d commit failed: %s", - batch_idx + 1, num_batches, e, + batch_idx + 1, + num_batches, + e, ) category = classify_error(e) if not _ErrorCategory.is_retryable(category): @@ -967,10 +920,11 @@ def _upload_worker( error_type="commit_" + category, ) logger.error( - "Batch %d/%d: permanent failure (%s), " - "%d file(s) will not be retried.", - batch_idx + 1, num_batches, - category, len(results), + "Batch %d/%d: permanent failure (%s), %d file(s) will not be retried.", + batch_idx + 1, + num_batches, + category, + len(results), ) consecutive_failures += 1 else: @@ -982,10 +936,11 @@ def _upload_worker( ) ) logger.warning( - "Batch %d/%d: %d file(s) recovered to retry " - "queue (error_category=%s).", - batch_idx + 1, num_batches, - len(results), category, + "Batch %d/%d: %d file(s) recovered to retry queue (error_category=%s).", + batch_idx + 1, + num_batches, + len(results), + category, ) consecutive_failures += 1 @@ -999,17 +954,15 @@ def _upload_worker( # ReAct progressive retry fallback if total_failed_files and UPLOAD_REACT_ENABLED: - total_failed_files, react_commits, react_results = ( - self._retry_failed_files_react( - failed_files=total_failed_files, - tracker=tracker, - repo_id=repo_id, - repo_type=repo_type, - commit_message=commit_message, - revision=revision, - max_workers=max_workers, - disable_tqdm=disable_tqdm, - ) + total_failed_files, react_commits, react_results = self._retry_failed_files_react( + failed_files=total_failed_files, + tracker=tracker, + repo_id=repo_id, + repo_type=repo_type, + commit_message=commit_message, + revision=revision, + max_workers=max_workers, + disable_tqdm=disable_tqdm, ) commit_infos.extend(react_commits) all_results.extend(react_results) @@ -1042,10 +995,7 @@ def _upload_worker( ) if orphans: delete_ops = self._build_delete_operations(orphans) - delete_commit_message = ( - f"{commit_message} " - f"(sync: delete {len(orphans)} orphan file(s))" - ) + delete_commit_message = f"{commit_message} (sync: delete {len(orphans)} orphan file(s))" try: delete_commit = self._commit_with_retry( repo_id=repo_id, @@ -1067,16 +1017,9 @@ def _upload_worker( elapsed = time.time() - start_time total_files = len(sorted_files) failed_count = len(total_failed_files) - lfs_reused_count = sum( - 1 for r in all_results - if r.get("upload_mode") == "lfs" and r.get("is_reused") - ) - lfs_uploaded_count = sum( - 1 for r in all_results if r.get("is_blob_uploaded") - ) - normal_count = sum( - 1 for r in all_results if r.get("upload_mode") == "normal" - ) + lfs_reused_count = sum(1 for r in all_results if r.get("upload_mode") == "lfs" and r.get("is_reused")) + lfs_uploaded_count = sum(1 for r in all_results if r.get("is_blob_uploaded")) + normal_count = sum(1 for r in all_results if r.get("upload_mode") == "normal") committed_count = len(all_results) print("=" * 60) @@ -1141,9 +1084,7 @@ def _compute_remote_orphans( item_type = item.get("Type") or item.get("type") or "blob" if item_type == "tree": continue - path = ( - item.get("Path") or item.get("path") or item.get("Name") or "" - ) + path = item.get("Path") or item.get("path") or item.get("Name") or "" if not path: continue remote_paths.append(path) @@ -1151,14 +1092,13 @@ def _compute_remote_orphans( # Filter by prefix scope if path_in_repo_prefix: scope_prefix = path_in_repo_prefix + "/" - remote_paths = [ - p for p in remote_paths if p.startswith(scope_prefix) - ] + remote_paths = [p for p in remote_paths if p.startswith(scope_prefix)] orphans = [p for p in remote_paths if p not in local_paths_in_repo] logger.info( "Sync: %d remote file(s) in scope, %d orphan(s) detected.", - len(remote_paths), len(orphans), + len(remote_paths), + len(orphans), ) return orphans @@ -1201,9 +1141,7 @@ def _upload_single_file( if is_real_path: try: file_stat = os.stat(file_path) - cached = tracker.get_hash( - file_path_in_repo, file_stat.st_mtime, file_stat.st_size - ) + cached = tracker.get_hash(file_path_in_repo, file_stat.st_mtime, file_stat.st_size) if cached is not None: hash_info_d = cached hash_info_d["file_path_or_obj"] = file_path @@ -1217,8 +1155,10 @@ def _upload_single_file( if file_stat is None: file_stat = os.stat(file_path) tracker.put_hash( - file_path_in_repo, file_stat.st_mtime, - file_stat.st_size, hash_info_d, + file_path_in_repo, + file_stat.st_mtime, + file_stat.st_size, + hash_info_d, ) except OSError: pass @@ -1253,10 +1193,7 @@ def _upload_single_file( sha256=file_hash, size=file_size, data=file_path, - disable_tqdm=( - disable_tqdm - or file_size <= UPLOAD_BLOB_TQDM_DISABLE_THRESHOLD - ), + disable_tqdm=(disable_tqdm or file_size <= UPLOAD_BLOB_TQDM_DISABLE_THRESHOLD), tqdm_desc=f"[Uploading {file_path_in_repo}]", pre_validated=pre_validated, ) @@ -1267,20 +1204,21 @@ def _upload_single_file( last_error = e if attempt < UPLOAD_BLOB_MAX_RETRIES - 1: wait = min( - UPLOAD_BLOB_RETRY_BACKOFF ** attempt, + UPLOAD_BLOB_RETRY_BACKOFF**attempt, UPLOAD_BLOB_RETRY_MAX_WAIT, ) logger.warning( - "Blob upload attempt %d/%d failed for %s: %s, " - "retrying in %ds ...", - attempt + 1, UPLOAD_BLOB_MAX_RETRIES, - file_path_in_repo, e, wait, + "Blob upload attempt %d/%d failed for %s: %s, retrying in %ds ...", + attempt + 1, + UPLOAD_BLOB_MAX_RETRIES, + file_path_in_repo, + e, + wait, ) time.sleep(wait) else: raise StorageError( - f"Blob upload failed after {UPLOAD_BLOB_MAX_RETRIES} attempts " - f"for {file_path_in_repo}: {last_error}" + f"Blob upload failed after {UPLOAD_BLOB_MAX_RETRIES} attempts for {file_path_in_repo}: {last_error}" ) from last_error else: if isinstance(file_path, (str, os.PathLike)): @@ -1302,10 +1240,7 @@ def _upload_single_file( "file_path_in_repo": file_path_in_repo, "file_path": file_path, "file_mtime": file_stat.st_mtime if file_stat else 0, - "file_size_on_disk": ( - file_stat.st_size if file_stat - else hash_info_d.get("file_size", 0) - ), + "file_size_on_disk": (file_stat.st_size if file_stat else hash_info_d.get("file_size", 0)), "is_uploaded": upload_res["is_uploaded"], "is_reused": upload_res.get("is_reused", False), "is_blob_uploaded": upload_res.get("is_blob_uploaded", False), @@ -1323,7 +1258,7 @@ def _upload_blob( repo_type: str, sha256: str, size: int, - data: str | Path | bytes | BinaryIO, + data: PathOrFileObj, disable_tqdm: bool = False, tqdm_desc: str = "[Uploading]", buffer_size_mb: int = 16, @@ -1343,21 +1278,20 @@ def _upload_blob( return res_d if isinstance(pre_validated, str): - upload_url = pre_validated + upload_url: str = pre_validated else: validated = self._client.validate_blobs( repo_id=repo_id, repo_type=repo_type, objects=[{"oid": sha256, "size": size}], ) - upload_url = validated.get(sha256) - if upload_url is None: - logger.info( - "Blob %s already exists globally, reuse.", sha256[:8] - ) + maybe_url = validated.get(sha256) + if maybe_url is None: + logger.info("Blob %s already exists globally, reuse.", sha256[:8]) res_d["is_uploaded"] = True res_d["is_reused"] = True return res_d + upload_url = maybe_url chunk_size = buffer_size_mb * 1024 * 1024 @@ -1371,23 +1305,15 @@ def _upload_blob( if isinstance(data, (str, Path)): with open(data, "rb") as f: stream = _CountedReadStream(f, size, pbar, chunk_size) - self._client.upload_blob( - upload_url=upload_url, data=stream, size=size - ) + self._client.upload_blob(upload_url=upload_url, data=stream, size=size) stream.verify_complete() elif isinstance(data, bytes): - stream = _CountedReadStream( - io.BytesIO(data), size, pbar, chunk_size - ) - self._client.upload_blob( - upload_url=upload_url, data=stream, size=size - ) + stream = _CountedReadStream(io.BytesIO(data), size, pbar, chunk_size) + self._client.upload_blob(upload_url=upload_url, data=stream, size=size) stream.verify_complete() else: stream = _CountedReadStream(data, size, pbar, chunk_size) - self._client.upload_blob( - upload_url=upload_url, data=stream, size=size - ) + self._client.upload_blob(upload_url=upload_url, data=stream, size=size) stream.verify_complete() res_d["url"] = upload_url @@ -1431,7 +1357,7 @@ def _commit_with_retry( revision: str = "master", max_retries: int = UPLOAD_COMMIT_MAX_RETRIES, ) -> dict: - last_error = None + last_error: Exception | None = None start_time = time.monotonic() for attempt in range(max_retries): try: @@ -1465,15 +1391,16 @@ def _commit_with_retry( break logger.warning( "Commit attempt %d/%d failed: %s, retrying in %ds ...", - attempt + 1, max_retries, last_error, wait, + attempt + 1, + max_retries, + last_error, + wait, ) time.sleep(wait) if isinstance(last_error, HubError): raise last_error - raise NetworkError( - f"Commit failed after {max_retries} attempts: {last_error}" - ) from last_error + raise NetworkError(f"Commit failed after {max_retries} attempts: {last_error}") from last_error # ------------------------------------------------------------------ # Internal: build operations @@ -1515,9 +1442,7 @@ def _build_operation( "encoding": "base64", } - def _build_batch_operations( - self, results: list[dict], repo_type: str - ) -> list[dict]: + def _build_batch_operations(self, results: list[dict], repo_type: str) -> list[dict]: operations = [] for item_d in results: file_path = item_d["file_path"] @@ -1546,9 +1471,7 @@ def _track_uploaded_batch( results: list[dict], ) -> None: for r in results: - tracker.mark_uploaded( - r["file_path_in_repo"], r["file_mtime"], r["file_size_on_disk"] - ) + tracker.mark_uploaded(r["file_path_in_repo"], r["file_mtime"], r["file_size_on_disk"]) tracker.save() def _track_committed_batch( @@ -1557,10 +1480,7 @@ def _track_committed_batch( results: list[dict], ) -> None: tracker.mark_committed_batch( - [ - (r["file_path_in_repo"], r["file_mtime"], r["file_size_on_disk"]) - for r in results - ] + [(r["file_path_in_repo"], r["file_mtime"], r["file_size_on_disk"]) for r in results] ) tracker.save() @@ -1579,15 +1499,10 @@ def _prepare_upload_folder( if not folder.is_dir(): raise InvalidParameter(f"Provided path: '{folder}' is not a directory") - all_files = sorted( - path for path in folder.glob("**/*") if path.is_file() - ) + all_files = sorted(path for path in folder.glob("**/*") if path.is_file()) if len(all_files) > UPLOAD_MAX_FILE_COUNT: - raise InvalidParameter( - f"Too many files ({len(all_files)}) in folder, " - f"max allowed: {UPLOAD_MAX_FILE_COUNT}" - ) + raise InvalidParameter(f"Too many files ({len(all_files)}) in folder, max allowed: {UPLOAD_MAX_FILE_COUNT}") # Per-directory file count check dir_counts: dict[str, int] = {} @@ -1623,10 +1538,7 @@ def _prepare_upload_folder( UPLOAD_NORMAL_FILE_SIZE_TOTAL_LIMIT, ) - relpath_to_abspath = { - path.relative_to(folder).as_posix(): str(path) - for path in all_files - } + relpath_to_abspath = {path.relative_to(folder).as_posix(): str(path) for path in all_files} filtered_keys = _filter_repo_objects( list(relpath_to_abspath.keys()), @@ -1635,10 +1547,7 @@ def _prepare_upload_folder( ) prefix = f"{path_in_repo.strip('/')}/" if path_in_repo else "" - prepared = [ - (prefix + relpath, relpath_to_abspath[relpath]) - for relpath in filtered_keys - ] + prepared = [(prefix + relpath, relpath_to_abspath[relpath]) for relpath in filtered_keys] logger.info("Prepared %d files for upload.", len(prepared)) return prepared @@ -1672,11 +1581,7 @@ def _retry_failed_files_react( else: permanent_failures.append(item_err) try: - st = ( - os.stat(file_path_r) - if isinstance(file_path_r, (str, os.PathLike)) - else None - ) + st = os.stat(file_path_r) if isinstance(file_path_r, (str, os.PathLike)) else None except OSError: st = None tracker.mark_failed( @@ -1687,11 +1592,13 @@ def _retry_failed_files_react( ) logger.error( "[ReAct] Permanent failure: %s (%s: %s)", - path_in_repo_r, category, err, + path_in_repo_r, + category, + err, ) retryable = remaining - round_configs = [ + round_configs: list[dict[str, Any]] = [ { "name": "Round 1 (parallel)", "parallel": True, @@ -1722,16 +1629,15 @@ def _retry_failed_files_react( round_name = cfg["name"] logger.info( "[ReAct] %s: retrying %d file(s) ...", - round_name, len(retryable), + round_name, + len(retryable), ) round_successes: list[dict] = [] round_failures: list[tuple] = [] if cfg["parallel"] and len(retryable) > 1: - with ThreadPoolExecutor( - max_workers=cfg["workers"] - ) as executor: + with ThreadPoolExecutor(max_workers=cfg["workers"]) as executor: future_map: dict = {} for (path_in_repo_r, file_path_r), _err in retryable: future = executor.submit( @@ -1750,27 +1656,20 @@ def _retry_failed_files_react( result = future.result() round_successes.append(result) except Exception as e: - round_failures.append( - ((path_in_repo_r, file_path_r), e) - ) + round_failures.append(((path_in_repo_r, file_path_r), e)) else: - for i, ((path_in_repo_r, file_path_r), _err) in enumerate( - retryable - ): + for i, ((path_in_repo_r, file_path_r), _err) in enumerate(retryable): if cfg["delay"] > 0 and i > 0: delay = ( - cfg["delay"] - * ( - 2 - ** min(i, UPLOAD_REACT_BACKOFF_MAX_EXPONENT) - ) + cfg["delay"] * (2 ** min(i, UPLOAD_REACT_BACKOFF_MAX_EXPONENT)) if round_idx == 1 else cfg["delay"] ) delay = min(delay, UPLOAD_REACT_MAX_DELAY) logger.info( "[ReAct] Waiting %ds before retrying %s ...", - delay, path_in_repo_r, + delay, + path_in_repo_r, ) time.sleep(delay) try: @@ -1786,17 +1685,15 @@ def _retry_failed_files_react( except Exception as e: logger.error( "[ReAct] %s: failed %s - %s", - round_name, path_in_repo_r, e, - ) - round_failures.append( - ((path_in_repo_r, file_path_r), e) + round_name, + path_in_repo_r, + e, ) + round_failures.append(((path_in_repo_r, file_path_r), e)) all_successes.extend(round_successes) - batch_size = min( - cfg["batch_size"], max(1, len(round_successes)) - ) + batch_size = min(cfg["batch_size"], max(1, len(round_successes))) for batch_start in range(0, len(round_successes), batch_size): batch = round_successes[batch_start : batch_start + batch_size] self._track_uploaded_batch(tracker, batch) @@ -1816,12 +1713,11 @@ def _retry_failed_files_react( self._track_committed_batch(tracker, batch) logger.info( "[ReAct] %s: committed %d file(s).", - round_name, len(batch), + round_name, + len(batch), ) except Exception as e: - logger.error( - "[ReAct] %s commit failed: %s", round_name, e - ) + logger.error("[ReAct] %s commit failed: %s", round_name, e) category = classify_error(e) if not _ErrorCategory.is_retryable(category): for r in batch: @@ -1833,24 +1729,16 @@ def _retry_failed_files_react( ) else: for r in batch: - round_failures.append( - ((r["file_path_in_repo"], r["file_path"]), e) - ) + round_failures.append(((r["file_path_in_repo"], r["file_path"]), e)) new_retryable = [] for item_err in round_failures: (path_in_repo_r, file_path_r), err = item_err - retry_counts[path_in_repo_r] = ( - retry_counts.get(path_in_repo_r, 0) + 1 - ) + retry_counts[path_in_repo_r] = retry_counts.get(path_in_repo_r, 0) + 1 if retry_counts[path_in_repo_r] >= 3: permanent_failures.append(item_err) try: - st = ( - os.stat(file_path_r) - if isinstance(file_path_r, (str, os.PathLike)) - else None - ) + st = os.stat(file_path_r) if isinstance(file_path_r, (str, os.PathLike)) else None except OSError: st = None tracker.mark_failed( @@ -1859,9 +1747,7 @@ def _retry_failed_files_react( st.st_size if st else 0, error_type="max_retries_exceeded", ) - logger.error( - "[ReAct] Max retries exceeded for %s", path_in_repo_r - ) + logger.error("[ReAct] Max retries exceeded for %s", path_in_repo_r) continue category = classify_error(err) if _ErrorCategory.is_retryable(category): @@ -1869,11 +1755,7 @@ def _retry_failed_files_react( else: permanent_failures.append(item_err) try: - st = ( - os.stat(file_path_r) - if isinstance(file_path_r, (str, os.PathLike)) - else None - ) + st = os.stat(file_path_r) if isinstance(file_path_r, (str, os.PathLike)) else None except OSError: st = None tracker.mark_failed( @@ -1884,15 +1766,17 @@ def _retry_failed_files_react( ) logger.error( "[ReAct] Permanent failure: %s (%s)", - path_in_repo_r, category, + path_in_repo_r, + category, ) progress = len(retryable) - len(new_retryable) if progress > 0: logger.info( - "[ReAct] %s: made progress — %d file(s) resolved, " - "%d remaining.", - round_name, progress, len(new_retryable), + "[ReAct] %s: made progress — %d file(s) resolved, %d remaining.", + round_name, + progress, + len(new_retryable), ) elif new_retryable: logger.warning( @@ -1933,7 +1817,8 @@ def _retry_failed_simple( break logger.info( "Retry round %d/%d: re-uploading %d failed file(s) ...", - retry_round + 1, UPLOAD_FAILED_FILE_MAX_RETRIES, + retry_round + 1, + UPLOAD_FAILED_FILE_MAX_RETRIES, len(total_failed_files), ) retry_failures: list[tuple] = [] @@ -1951,39 +1836,32 @@ def _retry_failed_simple( retry_successes.append(result) except Exception as e: logger.error(" Retry failed: %s - %s", path_in_repo_r, e) - retry_failures.append( - ((path_in_repo_r, file_path_r), e) - ) + retry_failures.append(((path_in_repo_r, file_path_r), e)) if retry_successes: self._track_uploaded_batch(tracker, retry_successes) - operations = self._build_batch_operations( - retry_successes, repo_type - ) + operations = self._build_batch_operations(retry_successes, repo_type) if operations: try: commit_info = self._commit_with_retry( repo_id=repo_id, repo_type=repo_type, operations=operations, - commit_message=( - f"{commit_message} " - f"(retry round {retry_round + 1})" - ), + commit_message=(f"{commit_message} (retry round {retry_round + 1})"), revision=revision, ) commit_infos.append(commit_info) all_results.extend(retry_successes) - self._track_committed_batch( - tracker, retry_successes - ) + self._track_committed_batch(tracker, retry_successes) logger.info( " Retry round %d: committed %d file(s).", - retry_round + 1, len(retry_successes), + retry_round + 1, + len(retry_successes), ) except Exception as e: logger.error( " Retry round %d commit failed: %s", - retry_round + 1, e, + retry_round + 1, + e, ) category = classify_error(e) if not _ErrorCategory.is_retryable(category): diff --git a/src/modelscope_hub/agent/__init__.py b/src/modelscope_hub/agent/__init__.py index 81d0cd9..50b788c 100644 --- a/src/modelscope_hub/agent/__init__.py +++ b/src/modelscope_hub/agent/__init__.py @@ -12,6 +12,7 @@ - :class:`RemoteFileInfo` -- metadata for a single remote file. - :func:`is_lfs_file` -- decide whether a file must use the LFS upload path. """ + from ._api import AgentApi, RemoteFileInfo, is_lfs_file __all__ = [ diff --git a/src/modelscope_hub/agent/_api.py b/src/modelscope_hub/agent/_api.py index 32c0965..2939ab7 100644 --- a/src/modelscope_hub/agent/_api.py +++ b/src/modelscope_hub/agent/_api.py @@ -12,6 +12,7 @@ * ``POST /api/v1/repos/agents/{id}/info/lfs/objects/batch`` -> LFS batch verify * ``DELETE /api/v1/agents/{path}/{name}/repo/file`` -> delete file """ + from __future__ import annotations import hashlib @@ -19,27 +20,70 @@ import os from dataclasses import dataclass -import requests - +from .._openapi import OpenAPIClient from ..config import HubConfig from ..constants import Visibility -from ..errors import APIError, AuthenticationError, HubError, NotExistError -from .._openapi import OpenAPIClient +from ..errors import AuthenticationError, NotExistError logger = logging.getLogger("modelscope_hub.agent") # LFS file extensions that must use LFS upload pathway. -_LFS_EXTENSIONS: frozenset[str] = frozenset({ - ".7z", ".aac", ".arrow", ".audio", ".bin", ".bmp", ".bz2", - ".ckpt", ".flac", ".ftz", ".gif", ".gz", ".h5", - ".jack", ".jpeg", ".jpg", ".joblib", ".jsonl", - ".lz4", ".mlmodel", ".model", ".mp3", ".mp4", ".msgpack", - ".npy", ".npz", ".ogg", ".onnx", ".ot", - ".parquet", ".pb", ".pcm", ".pickle", ".pkl", ".png", - ".pt", ".pth", ".rar", ".raw", - ".safetensors", ".sam", ".tar", ".tflite", ".tgz", ".tiff", - ".wasm", ".wav", ".webm", ".webp", ".xz", ".zip", ".zst", -}) +_LFS_EXTENSIONS: frozenset[str] = frozenset( + { + ".7z", + ".aac", + ".arrow", + ".audio", + ".bin", + ".bmp", + ".bz2", + ".ckpt", + ".flac", + ".ftz", + ".gif", + ".gz", + ".h5", + ".jack", + ".jpeg", + ".jpg", + ".joblib", + ".jsonl", + ".lz4", + ".mlmodel", + ".model", + ".mp3", + ".mp4", + ".msgpack", + ".npy", + ".npz", + ".ogg", + ".onnx", + ".ot", + ".parquet", + ".pb", + ".pcm", + ".pickle", + ".pkl", + ".png", + ".pt", + ".pth", + ".rar", + ".raw", + ".safetensors", + ".sam", + ".tar", + ".tflite", + ".tgz", + ".tiff", + ".wasm", + ".wav", + ".webm", + ".webp", + ".xz", + ".zip", + ".zst", + } +) # Files larger than this threshold (bytes) use LFS upload. _LFS_SIZE_THRESHOLD: int = 1 * 1024 * 1024 # 1 MB @@ -48,6 +92,7 @@ @dataclass class RemoteFileInfo: """Metadata for a single file in the remote repository.""" + path: str sha256: str is_lfs: bool = False @@ -114,15 +159,15 @@ def repo_info(self, path: str, name: str) -> dict | None: (e.g. the upload framework guard) consume. """ try: - return self._openapi.request( - "GET", f"/agents/{path}/{name}", require_token=False) + return self._openapi.request("GET", f"/agents/{path}/{name}", require_token=False) except NotExistError: return None except AuthenticationError: probe_url = f"{self.server}/api/v1/agents/{path}/{name}/repo/files" try: self._openapi.request( - "GET", url=probe_url, + "GET", + url=probe_url, params={"page_size": "1", "page": "1"}, require_token=False, ) @@ -143,11 +188,13 @@ def list_agents(self, owner: str | None = None, page_number: int = 1, page_size: """ criterion: list[dict] = [] if owner: - criterion.append({ - "Category": "Path", - "Predicate": "contains", - "StringValues": [owner], - }) + criterion.append( + { + "Category": "Path", + "Predicate": "contains", + "StringValues": [owner], + } + ) body = { "PageSize": page_size, "PageNumber": page_number, @@ -156,21 +203,18 @@ def list_agents(self, owner: str | None = None, page_number: int = 1, page_size: "Criterion": criterion, } list_url = f"{self.server}/api/v1/dolphin/agents" - data = self._openapi.request( - "PUT", url=list_url, json_body=body, require_token=False) + data = self._openapi.request("PUT", url=list_url, json_body=body, require_token=False) if isinstance(data, list): return {"items": data, "total_count": len(data)} if isinstance(data, dict): - items = next( - (data[k] for k in ("AgentList", "Agents", "agents", "Data", "data") - if k in data), + items: list = next( + (data[k] for k in ("AgentList", "Agents", "agents", "Data", "data") if k in data), [], ) if not isinstance(items, list): items = [] total_val = next( - (data[k] for k in ("TotalCount", "Total", "total_count") - if k in data and data[k] is not None), + (data[k] for k in ("TotalCount", "Total", "total_count") if k in data and data[k] is not None), len(items), ) try: @@ -180,8 +224,7 @@ def list_agents(self, owner: str | None = None, page_number: int = 1, page_size: return {"items": items, "total_count": total} return {"items": [], "total_count": 0} - def create_repo(self, path: str, name: str, framework: str | None = None, - visibility: str = "public") -> dict: + def create_repo(self, path: str, name: str, framework: str | None = None, visibility: str = "public") -> dict: """Create an empty agent (POST /agents). The server creates a bare repository. Files are added separately via @@ -196,30 +239,31 @@ def create_repo(self, path: str, name: str, framework: str | None = None, """ allowed = (Visibility.PUBLIC.label, Visibility.PRIVATE.label) if visibility not in allowed: - raise ValueError( - f"visibility must be one of {allowed}, got {visibility!r}") + raise ValueError(f"visibility must be one of {allowed}, got {visibility!r}") body: dict = {"path": path, "name": name, "visibility": visibility} if framework: body["framework"] = framework return self._openapi.request("POST", "/agents", json_body=body) - def list_repo_files(self, path: str, name: str, revision: str = 'master') -> list[str]: + def list_repo_files(self, path: str, name: str, revision: str = "master") -> list[str]: """All file paths in the repo, recursing into sub-directories.""" entries = self._fetch_tree_entries(path, name, revision) return [e["path"] for e in entries if e["type"] == "blob" and e["path"]] - def list_repo_files_detail(self, path: str, name: str, revision: str = 'master') -> list[RemoteFileInfo]: + def list_repo_files_detail(self, path: str, name: str, revision: str = "master") -> list[RemoteFileInfo]: """All blob files with sha256 and is_lfs flag.""" entries = self._fetch_tree_entries(path, name, revision) results: list[RemoteFileInfo] = [] for item in entries: if item["type"] != "blob" or not item["path"]: continue - results.append(RemoteFileInfo( - path=item["path"], - sha256=item.get("sha256") or "", - is_lfs=bool(item.get("is_lfs", False)), - )) + results.append( + RemoteFileInfo( + path=item["path"], + sha256=item.get("sha256") or "", + is_lfs=bool(item.get("is_lfs", False)), + ) + ) return results def _fetch_tree_entries(self, path: str, name: str, revision: str) -> list[dict]: @@ -232,7 +276,8 @@ def _fetch_tree_entries(self, path: str, name: str, revision: str) -> list[dict] list_url = f"{self.server}/api/v1/agents/{path}/{name}/repo/files" while True: data = self._openapi.request( - "GET", url=list_url, + "GET", + url=list_url, params={ "recursive": "true", "page_size": str(page_size), @@ -242,7 +287,7 @@ def _fetch_tree_entries(self, path: str, name: str, revision: str) -> list[dict] require_token=False, ) - raw = [] + raw: list = [] if isinstance(data, dict): raw = data.get("Trees") or data.get("trees") or [] elif isinstance(data, list): @@ -251,12 +296,14 @@ def _fetch_tree_entries(self, path: str, name: str, revision: str) -> list[dict] for item in raw: if not isinstance(item, dict): continue - all_entries.append({ - "path": item.get("Path") or item.get("path") or "", - "type": item.get("Type") or item.get("type") or "", - "sha256": item.get("Sha256") or item.get("sha256") or "", - "is_lfs": bool(item.get("IsLfs") or item.get("is_lfs") or False), - }) + all_entries.append( + { + "path": item.get("Path") or item.get("path") or "", + "type": item.get("Type") or item.get("type") or "", + "sha256": item.get("Sha256") or item.get("sha256") or "", + "is_lfs": bool(item.get("IsLfs") or item.get("is_lfs") or False), + } + ) if len(raw) < page_size: break @@ -264,27 +311,30 @@ def _fetch_tree_entries(self, path: str, name: str, revision: str) -> list[dict] if page > max_pages: logger.warning( "Pagination limit reached (%d pages) for %s/%s; results may be incomplete.", - max_pages, path, name, + max_pages, + path, + name, ) break return all_entries - def download_repo_file(self, path: str, name: str, file_path: str, - revision: str = "master", *, binary: bool = False): + def download_repo_file( + self, path: str, name: str, file_path: str, revision: str = "master", *, binary: bool = False + ): """Download one repo file. Returns bytes when *binary=True*, otherwise str. """ dl_url = f"{self.server}/agents/{path}/{name}/resolve/{revision}/{file_path}" - resp = self._openapi.request( - "GET", url=dl_url, unwrap=False, require_token=False) + resp = self._openapi.request("GET", url=dl_url, unwrap=False, require_token=False) return resp.content if binary else resp.text # ---- commit (normal + LFS) ---- - def commit_files(self, path: str, name: str, actions: list[dict], - revision: str = "master", commit_message: str = "sync") -> dict: + def commit_files( + self, path: str, name: str, actions: list[dict], revision: str = "master", commit_message: str = "sync" + ) -> dict: """Commit file changes via POST /api/v1/repos/agents/{path}/{name}/commit/{revision}. Each action dict should contain: @@ -306,10 +356,7 @@ def lfs_batch(self, path: str, name: str, oid: str, size: int) -> str | None: POST /api/v1/repos/agents/{path}/{name}/info/lfs/objects/batch Returns the upload href if the server needs the blob, None otherwise. """ - batch_url = ( - f"{self.server}/api/v1/repos/agents/{path}/{name}" - f"/info/lfs/objects/batch" - ) + batch_url = f"{self.server}/api/v1/repos/agents/{path}/{name}/info/lfs/objects/batch" body = { "operation": "upload", "objects": [{"oid": oid, "size": size}], @@ -317,7 +364,7 @@ def lfs_batch(self, path: str, name: str, oid: str, size: int) -> str | None: data = self._openapi.request("POST", url=batch_url, json_body=body) # Response: {"objects": [{"actions": {"upload": {"href": ...}}}]} # If no actions.upload -> blob already exists, skip PUT. - objects = [] + objects: list = [] if isinstance(data, dict): objects = data.get("objects") or [] if not objects: @@ -328,7 +375,8 @@ def lfs_batch(self, path: str, name: str, oid: str, size: int) -> str | None: def lfs_upload_blob(self, upload_url: str, data: bytes) -> None: """PUT binary data to the LFS upload URL.""" self._openapi.request( - "PUT", url=upload_url, + "PUT", + url=upload_url, data=data, headers={"Content-Type": "application/octet-stream"}, require_token=False, @@ -336,10 +384,16 @@ def lfs_upload_blob(self, upload_url: str, data: bytes) -> None: timeout=max(self.timeout, 300), ) - def upload_lfs_file(self, path: str, name: str, file_path: str, - content: bytes, action: str = "create", - revision: str = "master", - commit_message: str = "sync") -> dict: + def upload_lfs_file( + self, + path: str, + name: str, + file_path: str, + content: bytes, + action: str = "create", + revision: str = "master", + commit_message: str = "sync", + ) -> dict: """Full LFS upload flow: batch verify -> PUT blob -> commit reference. Combines lfs_batch + lfs_upload_blob + commit_files for one file. @@ -354,21 +408,22 @@ def upload_lfs_file(self, path: str, name: str, file_path: str, self.lfs_upload_blob(upload_url, content) # Step 3: commit LFS reference - actions = [{ - "action": action, - "path": file_path, - "type": "lfs", - "size": size, - "sha256": oid, - "content": "", - "encoding": "", - }] - return self.commit_files(path, name, actions, revision=revision, - commit_message=commit_message) - - def delete_file(self, path: str, name: str, file_path: str, - revision: str = "master", - commit_message: str | None = None) -> dict: + actions = [ + { + "action": action, + "path": file_path, + "type": "lfs", + "size": size, + "sha256": oid, + "content": "", + "encoding": "", + } + ] + return self.commit_files(path, name, actions, revision=revision, commit_message=commit_message) + + def delete_file( + self, path: str, name: str, file_path: str, revision: str = "master", commit_message: str | None = None + ) -> dict: """Delete a file from the repo. DELETE /api/v1/agents/{path}/{name}/repo/file diff --git a/src/modelscope_hub/api.py b/src/modelscope_hub/api.py index 2907d35..3ce9a35 100644 --- a/src/modelscope_hub/api.py +++ b/src/modelscope_hub/api.py @@ -23,23 +23,23 @@ from __future__ import annotations +from collections.abc import Iterable, Mapping from pathlib import Path -from typing import Any, BinaryIO, Iterable, Mapping +from typing import Any, BinaryIO, TypeAlias from urllib.parse import urlparse from requests.cookies import RequestsCookieJar -from ._cache_manager import clear_cache as _clear_cache from ._cache_manager import _resolve_verification_root +from ._cache_manager import clear_cache as _clear_cache from ._cache_manager import scan_cache as _scan_cache -from ._download import DownloadManager, ProgressCallback from ._cache_manager import verify_cache as _verify_cache -from ._download import DownloadManager +from ._download import DownloadManager, ProgressCallback from ._legacy_api import LegacyClient from ._openapi import OpenAPIClient from ._upload import UploadManager from .config import HubConfig, get_default_config -from .constants import RepoType, Visibility +from .constants import DEFAULT_ENDPOINT, RepoType, Visibility from .errors import ( AuthenticationError, HubError, @@ -55,7 +55,7 @@ logger = get_logger("api") -RepoTypeLike = "str | RepoType" +RepoTypeLike: TypeAlias = "str | RepoType" # Routing tables — declarative dispatch keeps :class:`HubApi` free of long @@ -185,7 +185,7 @@ def legacy(self) -> LegacyClient: self._legacy = LegacyClient( token=self._config.token, - endpoint=self._config.endpoint, + endpoint=self._config.endpoint or DEFAULT_ENDPOINT, user_agent=build_user_agent(self._config.get_session_id()), ) elif self._legacy.token != self._config.token and self._config.token: @@ -999,7 +999,7 @@ def resolve_endpoint_for_read( ) if self._config._endpoint_overridden: - return self._config.endpoint + return self._config.endpoint or DEFAULT_ENDPOINT effective_token = token or self._config.token @@ -1485,7 +1485,7 @@ def delete_files( revision=revision or "master", ) deleted.append(p) - except (AuthenticationError, NetworkError) as exc: + except (AuthenticationError, NetworkError): failed.append(p) raise except Exception: diff --git a/src/modelscope_hub/cli/agent.py b/src/modelscope_hub/cli/agent.py index 5454da3..e99e0b3 100644 --- a/src/modelscope_hub/cli/agent.py +++ b/src/modelscope_hub/cli/agent.py @@ -11,13 +11,13 @@ import base64 import sys -from argparse import Action, RawDescriptionHelpFormatter +from argparse import RawDescriptionHelpFormatter from pathlib import Path from ..agent import AgentApi, is_lfs_file from ..constants import Visibility from ..errors import APIError -from .base import CLICommand +from .base import CLICommand, SubParsers _CONVERT_HINT = ( "This command transfers raw files only. For framework-aware conversion, " @@ -68,8 +68,7 @@ def _cmd_list(owner, page_number, page_size, *, endpoint, token) -> int: return _fail("not logged in. Provide endpoint.") client = AgentApi(endpoint=endpoint, token=token or "") try: - result = client.list_agents( - owner=owner, page_number=page_number, page_size=page_size) + result = client.list_agents(owner=owner, page_number=page_number, page_size=page_size) except APIError as e: return _fail(_api_error_message(e, "list")) except Exception as e: @@ -115,8 +114,8 @@ def _cmd_download(repo, local_dir, revision, *, endpoint, token, username) -> in return _fail("not logged in. Provide endpoint.") if "/" not in repo and not username: return _fail( - f"--repo '{repo}' requires login to resolve owner. " - f"Use 'owner/name' format or run 'ms login' first.") + f"--repo '{repo}' requires login to resolve owner. Use 'owner/name' format or run 'ms login' first." + ) group, name = _resolve_repo(repo, username or "") client = AgentApi(endpoint=endpoint, token=token or "") @@ -137,8 +136,7 @@ def _cmd_download(repo, local_dir, revision, *, endpoint, token, username) -> in for i, rel in enumerate(paths, 1): print(f" [{i}/{total}] downloading {rel}", flush=True) try: - data = client.download_repo_file( - group, name, rel, revision=revision, binary=True) + data = client.download_repo_file(group, name, rel, revision=revision, binary=True) except APIError as e: return _fail(_api_error_message(e, "download")) except Exception as e: @@ -150,8 +148,7 @@ def _cmd_download(repo, local_dir, revision, *, endpoint, token, username) -> in return 0 -def _cmd_upload(repo, local_dir, revision, dry_run, *, endpoint, token, username, - visibility="public") -> int: +def _cmd_upload(repo, local_dir, revision, dry_run, *, endpoint, token, username, visibility="public") -> int: """Upload raw files from a local path to a remote repository.""" if not repo: return _fail("--repo is required (the remote repository name).") @@ -182,8 +179,8 @@ def _cmd_upload(repo, local_dir, revision, dry_run, *, endpoint, token, username return _fail("not logged in. Run 'ms login' first.") if "/" not in repo and not username: return _fail( - f"--repo '{repo}' requires login to resolve owner. " - f"Use 'owner/name' format or run 'ms login' first.") + f"--repo '{repo}' requires login to resolve owner. Use 'owner/name' format or run 'ms login' first." + ) group, name = _resolve_repo(repo, username or "") client = AgentApi(endpoint=endpoint, token=token) @@ -191,8 +188,7 @@ def _cmd_upload(repo, local_dir, revision, dry_run, *, endpoint, token, username if not client.check_repo(group, name): client.create_repo(group, name, visibility=visibility) except Exception as exc: - print(f"warning: create_repo check failed ({exc}), proceeding anyway.", - file=sys.stderr) + print(f"warning: create_repo check failed ({exc}), proceeding anyway.", file=sys.stderr) # Normal files (< LFS threshold, non-LFS extension) are small by definition # and go in a single commit; LFS files are read one at a time to bound @@ -204,24 +200,30 @@ def _cmd_upload(repo, local_dir, revision, dry_run, *, endpoint, token, username if is_lfs_file(rel, size): lfs_entries.append((rel, fp)) else: - normal_actions.append({ - "action": "create", - "path": rel, - "type": "normal", - "size": size, - "sha256": "", - "content": base64.b64encode(fp.read_bytes()).decode("ascii"), - "encoding": "base64", - }) + normal_actions.append( + { + "action": "create", + "path": rel, + "type": "normal", + "size": size, + "sha256": "", + "content": base64.b64encode(fp.read_bytes()).decode("ascii"), + "encoding": "base64", + } + ) try: if normal_actions: - client.commit_files( - group, name, normal_actions, revision=revision, - commit_message="upload normal files") + client.commit_files(group, name, normal_actions, revision=revision, commit_message="upload normal files") for rel, fp in lfs_entries: client.upload_lfs_file( - group, name, rel, fp.read_bytes(), action="create", - revision=revision, commit_message=f"upload LFS {rel}") + group, + name, + rel, + fp.read_bytes(), + action="create", + revision=revision, + commit_message=f"upload LFS {rel}", + ) except APIError as e: return _fail(_api_error_message(e, "upload")) except Exception as e: @@ -238,7 +240,7 @@ class AgentCommand(CLICommand): """Raw agent-repository file transfer: download, upload, list.""" @staticmethod - def register(subparsers: Action) -> None: + def register(subparsers: SubParsers) -> None: _epilog = ( "subcommand arguments:\n" " download -r REPO [--local-dir DIR] [--revision REV]\n" @@ -256,8 +258,7 @@ def register(subparsers: Action) -> None: agent_parser = subparsers.add_parser( "agent", help="Transfer raw agent repository files (download, upload, list).", - description="Low-level raw file transfer for remote agent repositories. " - + _CONVERT_HINT, + description="Low-level raw file transfer for remote agent repositories. " + _CONVERT_HINT, epilog=_epilog, formatter_class=RawDescriptionHelpFormatter, ) @@ -270,42 +271,45 @@ def register(subparsers: Action) -> None: "download", help="Download raw agent files from a remote repository", formatter_class=RawDescriptionHelpFormatter, - description="Download all files of a remote agent repository to a local directory.\n" - + _CONVERT_HINT, + description="Download all files of a remote agent repository to a local directory.\n" + _CONVERT_HINT, ) p_download.add_argument( - "-r", "--repo", required=True, - help="Remote repo identifier, supports owner/name format (e.g. user/my-agent)") - p_download.add_argument( - "--local-dir", default=None, - help="Destination directory (default: ./ under CWD)") + "-r", + "--repo", + required=True, + help="Remote repo identifier, supports owner/name format (e.g. user/my-agent)", + ) p_download.add_argument( - "--revision", default="master", help="Repository revision (default: master)") + "--local-dir", default=None, help="Destination directory (default: ./ under CWD)" + ) + p_download.add_argument("--revision", default="master", help="Repository revision (default: master)") # ---- upload ---- p_upload = agent_sub.add_parser( "upload", help="Upload raw agent files to a remote repository", formatter_class=RawDescriptionHelpFormatter, - description="Upload files from a local path to a remote agent repository.\n" - + _CONVERT_HINT, + description="Upload files from a local path to a remote agent repository.\n" + _CONVERT_HINT, ) p_upload.add_argument( - "-r", "--repo", required=True, - help="Remote repo identifier, supports owner/name format (e.g. user/my-agent)") - p_upload.add_argument( - "--local-dir", default=None, - help="Source path (file or directory) to upload (default: CWD)") + "-r", + "--repo", + required=True, + help="Remote repo identifier, supports owner/name format (e.g. user/my-agent)", + ) p_upload.add_argument( - "--revision", default="master", help="Repository revision (default: master)") + "--local-dir", default=None, help="Source path (file or directory) to upload (default: CWD)" + ) + p_upload.add_argument("--revision", default="master", help="Repository revision (default: master)") p_upload.add_argument( "--visibility", choices=[Visibility.PUBLIC.label, Visibility.PRIVATE.label], default=Visibility.PUBLIC.label, - help="Visibility of the remote repo when created (default: public)") + help="Visibility of the remote repo when created (default: public)", + ) p_upload.add_argument( - "--dry-run", action="store_true", - help="List files that would be uploaded, without actually uploading") + "--dry-run", action="store_true", help="List files that would be uploaded, without actually uploading" + ) # ---- list ---- p_list = agent_sub.add_parser( @@ -313,21 +317,20 @@ def register(subparsers: Action) -> None: help="List remote agent repositories", description="Query and display remote agent repositories with pagination.", ) + p_list.add_argument("--owner", default=None, help="Filter by owner username or organization name") p_list.add_argument( - "--owner", default=None, - help="Filter by owner username or organization name") - p_list.add_argument( - "--page", dest="page_number", type=int, default=1, - help="Page number for pagination (default: 1)") + "--page", dest="page_number", type=int, default=1, help="Page number for pagination (default: 1)" + ) p_list.add_argument( - "--page-size", dest="page_size", type=int, default=10, - help="Number of items per page (default: 10)") + "--page-size", dest="page_size", type=int, default=10, help="Number of items per page (default: 10)" + ) def execute(self) -> None: args = self.args action = args.agent_command from ..config import HubConfig + config = HubConfig( endpoint=getattr(args, "endpoint", None), token=getattr(args, "token", None), @@ -337,12 +340,10 @@ def execute(self) -> None: # Resolve current username for repos given without an explicit owner. username = "" - needs_user = ( - action == "upload" - or (action == "download" and "/" not in getattr(args, "repo", "")) - ) + needs_user = action == "upload" or (action == "download" and "/" not in getattr(args, "repo", "")) if needs_user and token: from .._openapi import OpenAPIClient + try: openapi = OpenAPIClient(config=config) user_data = openapi.get_current_user() or {} diff --git a/src/modelscope_hub/cli/base.py b/src/modelscope_hub/cli/base.py index 630e36f..f2de5f6 100644 --- a/src/modelscope_hub/cli/base.py +++ b/src/modelscope_hub/cli/base.py @@ -14,14 +14,29 @@ import sys from abc import ABC, abstractmethod -from argparse import Action, ArgumentParser, Namespace -from typing import Any, Iterable, Sequence +from argparse import ArgumentParser, Namespace +from collections.abc import Iterable, Sequence +from typing import Any, Protocol from ..api import HubApi from ..constants import RepoType from ..utils.format import tabulate as _tabulate +# --------------------------------------------------------------------------- +# Structural type for argparse sub-parser containers +# --------------------------------------------------------------------------- +class SubParsers(Protocol): + """Structural stand-in for ``argparse._SubParsersAction``. + + argparse only exposes its sub-parsers container as a private class, so + commands accept anything that provides ``add_parser`` instead. This keeps + both mypy and IDE inspections happy without referencing private names. + """ + + def add_parser(self, name: str, **kwargs: Any) -> ArgumentParser: ... + + # --------------------------------------------------------------------------- # Abstract base class # --------------------------------------------------------------------------- @@ -33,7 +48,7 @@ def __init__(self, args: Namespace) -> None: @staticmethod @abstractmethod - def register(subparsers: Action) -> None: + def register(subparsers: SubParsers) -> None: """Attach this command's argparse parser to ``subparsers``.""" @abstractmethod @@ -78,7 +93,8 @@ def add_repo_type_arg( """ valid = list(choices) if choices else [t.value for t in RepoType] parser.add_argument( - "--repo-type", "--repo_type", + "--repo-type", + "--repo_type", dest="repo_type", choices=valid, default=default, @@ -124,9 +140,7 @@ def parse_kv_pairs(values: Iterable[str]) -> dict[str, str]: result: dict[str, str] = {} for raw in values: if "=" not in raw: - raise ValueError( - f"Invalid setting {raw!r}: expected 'key=value' format." - ) + raise ValueError(f"Invalid setting {raw!r}: expected 'key=value' format.") key, _, value = raw.partition("=") key = key.strip() if not key: diff --git a/src/modelscope_hub/cli/cache.py b/src/modelscope_hub/cli/cache.py index a78a1b1..d0372b7 100644 --- a/src/modelscope_hub/cli/cache.py +++ b/src/modelscope_hub/cli/cache.py @@ -3,18 +3,17 @@ from __future__ import annotations import sys -from argparse import Action from ..constants import RepoType from ..utils.format import format_size -from .base import CLICommand, error, info, make_api, render_table, success, warn +from .base import CLICommand, SubParsers, error, info, make_api, render_table, success, warn class CacheCommand(CLICommand): """Top-level dispatcher for the ``cache`` subcommands.""" @staticmethod - def register(subparsers: Action) -> None: + def register(subparsers: SubParsers) -> None: parser = subparsers.add_parser("cache", help="Inspect or clear the local cache.") sub = parser.add_subparsers(dest="cache_action", metavar="ACTION") sub.required = True @@ -39,7 +38,7 @@ def _human_size(num: int) -> str: class _CacheScan(CLICommand): @staticmethod - def register(subparsers: Action) -> None: + def register(subparsers: SubParsers) -> None: p = subparsers.add_parser("scan", help="Show cached repositories and disk usage.") p.add_argument("--cache-dir", dest="cache_dir", default=None) p.set_defaults(_command=CacheCommand, _cache_leaf=_CacheScan) @@ -73,7 +72,7 @@ def execute(self) -> None: class _CacheClear(CLICommand): @staticmethod - def register(subparsers: Action) -> None: + def register(subparsers: SubParsers) -> None: p = subparsers.add_parser("clear", help="Remove cached files.") p.add_argument("--cache-dir", dest="cache_dir", default=None) p.add_argument( @@ -119,7 +118,7 @@ def _format_paths(paths: list[str], limit: int = 10) -> str: class _CacheVerify(CLICommand): @staticmethod - def register(subparsers: Action) -> None: + def register(subparsers: SubParsers) -> None: p = subparsers.add_parser("verify", help="Verify local files against Hub SHA-256 checksums.") p.add_argument("repo_id", help="Repository id in owner/name form.") p.add_argument( diff --git a/src/modelscope_hub/cli/compat.py b/src/modelscope_hub/cli/compat.py index 47c60d7..0bf5920 100644 --- a/src/modelscope_hub/cli/compat.py +++ b/src/modelscope_hub/cli/compat.py @@ -10,8 +10,8 @@ import os import warnings from argparse import SUPPRESS, Action, ArgumentParser, Namespace -from typing import Any, Sequence - +from collections.abc import Sequence +from typing import Any # --------------------------------------------------------------------------- # Deprecation infrastructure @@ -25,8 +25,7 @@ def deprecated_arg(old: str, new: str) -> None: if os.environ.get(_SUPPRESS_ENVVAR) or os.environ.get(_SUPPRESS_ENVVAR_OLD): return warnings.warn( - f"'{old}' is deprecated and will be removed in a future version. " - f"Use '{new}' instead.", + f"'{old}' is deprecated and will be removed in a future version. Use '{new}' instead.", DeprecationWarning, stacklevel=3, ) @@ -49,9 +48,11 @@ def __call__( self, parser: ArgumentParser, namespace: Namespace, - values: str | Sequence[str], + values: str | Sequence[Any] | None, option_string: str | None = None, ) -> None: + if values is None: + return current: list[str] = getattr(namespace, self.dest, None) or [] if isinstance(values, str): current.append(values) @@ -140,10 +141,7 @@ def normalize_download_args(args: Namespace) -> Namespace: args.files = [] if not args.repo_id: - raise ValueError( - "repo_id is required. Provide it as a positional argument " - "or via --model/--dataset." - ) + raise ValueError("repo_id is required. Provide it as a positional argument or via --model/--dataset.") return args diff --git a/src/modelscope_hub/cli/deploy.py b/src/modelscope_hub/cli/deploy.py index 6faa904..8439f45 100644 --- a/src/modelscope_hub/cli/deploy.py +++ b/src/modelscope_hub/cli/deploy.py @@ -7,11 +7,11 @@ from __future__ import annotations import json -from argparse import Action from ..constants import RepoType from .base import ( CLICommand, + SubParsers, add_repo_type_arg, info, make_api, @@ -25,7 +25,7 @@ class DeployCommand(CLICommand): """Deploy a Studio space or an MCP server.""" @staticmethod - def register(subparsers: Action) -> None: + def register(subparsers: SubParsers) -> None: p = subparsers.add_parser("deploy", help="Deploy a studio space or MCP server.") p.add_argument("repo_id") add_repo_type_arg( @@ -47,7 +47,7 @@ class StopCommand(CLICommand): """Stop a running Studio or undeploy an MCP server.""" @staticmethod - def register(subparsers: Action) -> None: + def register(subparsers: SubParsers) -> None: p = subparsers.add_parser("stop", help="Stop a studio space or undeploy MCP server.") p.add_argument("repo_id") add_repo_type_arg( @@ -69,7 +69,7 @@ class LogsCommand(CLICommand): """Stream paginated run / build logs of a Studio space.""" @staticmethod - def register(subparsers: Action) -> None: + def register(subparsers: SubParsers) -> None: p = subparsers.add_parser("logs", help="Fetch logs for a studio space.") p.add_argument("repo_id") add_repo_type_arg( @@ -108,7 +108,7 @@ class SettingsCommand(CLICommand): """Update Studio / Skill settings via ``key=value`` tokens.""" @staticmethod - def register(subparsers: Action) -> None: + def register(subparsers: SubParsers) -> None: p = subparsers.add_parser( "settings", help="Update studio or skill settings (key=value pairs).", @@ -142,7 +142,10 @@ def _extract_log_lines(payload: object) -> list[str]: for key in ("logs", "items", "list", "data"): value = payload.get(key) if isinstance(value, list): - return [str(v) if not isinstance(v, dict) else (v.get("message") or json.dumps(v, ensure_ascii=False)) for v in value] + return [ + str(v) if not isinstance(v, dict) else (v.get("message") or json.dumps(v, ensure_ascii=False)) + for v in value + ] if isinstance(value, str): return value.splitlines() return [] diff --git a/src/modelscope_hub/cli/download.py b/src/modelscope_hub/cli/download.py index b9bf196..328fc17 100644 --- a/src/modelscope_hub/cli/download.py +++ b/src/modelscope_hub/cli/download.py @@ -3,13 +3,12 @@ from __future__ import annotations import sys -from argparse import Action from concurrent.futures import ThreadPoolExecutor, as_completed from pathlib import Path from ..api import HubApi from ..constants import RepoType -from .base import CLICommand, add_repo_type_arg, info, make_api, success, warn +from .base import CLICommand, SubParsers, add_repo_type_arg, info, make_api, success, warn from .compat import ( PatternAction, add_legacy_download_args, @@ -35,7 +34,7 @@ class DownloadCommand(CLICommand): """Download files or whole repositories from ModelScope Hub.""" @staticmethod - def register(subparsers: Action) -> None: + def register(subparsers: SubParsers) -> None: p = subparsers.add_parser( "download", help="Download a file or full snapshot of a repository.", @@ -59,8 +58,9 @@ def register(subparsers: Action) -> None: ) p.add_argument("--revision", default=None, help="Branch / tag / commit (default: master).") p.add_argument("--cache-dir", dest="cache_dir", default=None, help="Override cache directory.") - p.add_argument("--local-dir", dest="local_dir", default=None, - help="Download directly to this directory (bypasses cache).") + p.add_argument( + "--local-dir", dest="local_dir", default=None, help="Download directly to this directory (bypasses cache)." + ) p.add_argument( "--max-workers", dest="max_workers", @@ -151,7 +151,8 @@ def _make_api_with_endpoint(self) -> HubApi: api = make_api(self.args) try: resolved = api.resolve_endpoint_for_read( - self.args.repo_id, repo_type=self.args.repo_type, + self.args.repo_id, + repo_type=self.args.repo_type, ) return HubApi(token=token, endpoint=resolved) except Exception: @@ -164,13 +165,8 @@ def _download_collection(self) -> None: collection_id = self.args.repo_id data = api.legacy.get_collection(collection_id) - elements = data.get("CollectionElements", {}).get( - "CollectionElementVoList", [] - ) - valid = [ - e for e in elements - if e.get("ElementPath") and e.get("ElementName") - ] + elements = data.get("CollectionElements", {}).get("CollectionElementVoList", []) + valid = [e for e in elements if e.get("ElementPath") and e.get("ElementName")] if not valid: warn(f"No valid skill elements found in collection: {collection_id}") return @@ -203,10 +199,7 @@ def _download_one(elem: dict) -> tuple[str, str | None, str | None]: succeeded.append((sid, path)) success(f"skill {sid} → {path}") - info( - f"Download complete: {len(succeeded)} succeeded, " - f"{len(failed)} failed" - ) + info(f"Download complete: {len(succeeded)} succeeded, {len(failed)} failed") if failed: for sid, err in failed: warn(f" {sid}: {err}") diff --git a/src/modelscope_hub/cli/login.py b/src/modelscope_hub/cli/login.py index b1974d0..d7500da 100644 --- a/src/modelscope_hub/cli/login.py +++ b/src/modelscope_hub/cli/login.py @@ -8,9 +8,9 @@ from __future__ import annotations import getpass -from argparse import Action, SUPPRESS +from argparse import SUPPRESS -from .base import CLICommand, error, info, make_api, success +from .base import CLICommand, SubParsers, error, info, make_api, success from .compat import add_subcmd_token_endpoint @@ -18,7 +18,7 @@ class LoginCommand(CLICommand): """Persist a token and verify it via ``GET /users/me``.""" @staticmethod - def register(subparsers: Action) -> None: + def register(subparsers: SubParsers) -> None: parser = subparsers.add_parser( "login", help="Authenticate with ModelScope Hub and persist the token locally.", @@ -61,7 +61,7 @@ class WhoamiCommand(CLICommand): """Show the currently authenticated user.""" @staticmethod - def register(subparsers: Action) -> None: + def register(subparsers: SubParsers) -> None: parser = subparsers.add_parser( "whoami", help="Show the user identified by the active token.", diff --git a/src/modelscope_hub/cli/main.py b/src/modelscope_hub/cli/main.py index f46a7e4..6fa219d 100644 --- a/src/modelscope_hub/cli/main.py +++ b/src/modelscope_hub/cli/main.py @@ -19,13 +19,13 @@ import logging import sys from argparse import SUPPRESS -from typing import Sequence +from collections.abc import Sequence from .. import __version__ from ..constants import MODELSCOPE_ASCII -from ..errors import HubError, InvalidParameter, NetworkError, NotSupportedError -from .base import CLICommand, add_repo_type_arg, error, info, make_api, success +from ..errors import HubError, InvalidParameter, NotSupportedError from .agent import AgentCommand +from .base import CLICommand, error, info from .cache import CacheCommand, _CacheClear, _CacheScan from .deploy import DeployCommand, LogsCommand, SettingsCommand, StopCommand from .download import DownloadCommand @@ -86,7 +86,8 @@ def _build_parser() -> argparse.ArgumentParser: help="API endpoint (overrides MODELSCOPE_ENDPOINT).", ) parser.add_argument( - "-v", "--verbose", + "-v", + "--verbose", action="store_true", help="Enable verbose (DEBUG) logging.", ) @@ -125,7 +126,6 @@ def _register_scan_cache_alias(subparsers) -> None: def _register_clear_cache_alias(subparsers) -> None: """``ms-hub clear-cache`` → alias for ``ms-hub cache clear``.""" - from ..constants import RepoType p = subparsers.add_parser("clear-cache", help="[Alias] Remove cached files.") group = p.add_mutually_exclusive_group() @@ -136,7 +136,6 @@ def _register_clear_cache_alias(subparsers) -> None: p.set_defaults(_command=_ClearCacheAlias) - class _ScanCacheAlias(CLICommand): """Adapter: top-level ``scan-cache`` → ``cache scan``.""" @@ -178,10 +177,8 @@ def execute(self) -> None: # --------------------------------------------------------------------------- def _discover_plugins(subparsers) -> None: """Discover CLI plugins registered via entry_points.""" - try: - eps = importlib.metadata.entry_points(group=_PLUGIN_GROUP) - except TypeError: - eps = importlib.metadata.entry_points().get(_PLUGIN_GROUP, []) + # ``entry_points(group=...)`` is available on all supported Pythons (3.10+). + eps = importlib.metadata.entry_points(group=_PLUGIN_GROUP) for ep in eps: try: @@ -191,9 +188,7 @@ def _discover_plugins(subparsers) -> None: elif hasattr(cmd_cls, "define_args"): cmd_cls.define_args(subparsers) except Exception as exc: - logging.getLogger(__name__).debug( - "Failed to load CLI plugin %r: %s", ep.name, exc - ) + logging.getLogger(__name__).debug("Failed to load CLI plugin %r: %s", ep.name, exc) # --------------------------------------------------------------------------- diff --git a/src/modelscope_hub/cli/mcp.py b/src/modelscope_hub/cli/mcp.py index b3f8d9d..d383dbf 100644 --- a/src/modelscope_hub/cli/mcp.py +++ b/src/modelscope_hub/cli/mcp.py @@ -3,9 +3,8 @@ from __future__ import annotations import json -from argparse import Action -from .base import CLICommand, info, make_api, render_table, success +from .base import CLICommand, SubParsers, info, make_api, parse_kv_pairs, render_table, success from .compat import add_subcmd_token_endpoint @@ -13,7 +12,7 @@ class McpCommand(CLICommand): """Top-level dispatcher for the ``mcp`` subcommands.""" @staticmethod - def register(subparsers: Action) -> None: + def register(subparsers: SubParsers) -> None: parser = subparsers.add_parser("mcp", help="Manage MCP servers.") sub = parser.add_subparsers(dest="mcp_action", metavar="ACTION") sub.required = True @@ -34,7 +33,7 @@ def execute(self) -> None: class _McpList(CLICommand): @staticmethod - def register(subparsers: Action) -> None: + def register(subparsers: SubParsers) -> None: p = subparsers.add_parser("list", help="List MCP servers.") p.add_argument("--search", default=None) p.add_argument("--page", dest="page_number", type=int, default=1) @@ -67,7 +66,7 @@ def execute(self) -> None: class _McpInfo(CLICommand): @staticmethod - def register(subparsers: Action) -> None: + def register(subparsers: SubParsers) -> None: p = subparsers.add_parser("info", help="Show details of an MCP server.") p.add_argument("server_id") add_subcmd_token_endpoint(p) @@ -81,16 +80,38 @@ def execute(self) -> None: class _McpDeploy(CLICommand): @staticmethod - def register(subparsers: Action) -> None: + def register(subparsers: SubParsers) -> None: p = subparsers.add_parser("deploy", help="Deploy an MCP server.") p.add_argument("server_id") p.add_argument( - "--transport-type", dest="transport_type", default=None, + "--transport-type", + dest="transport_type", + default=None, + choices=["sse", "streamable_http"], help="Transport type (default: sse).", ) p.add_argument( - "--expiration-minutes", dest="expiration_minutes", type=int, default=None, - help="Expiration time in minutes.", + "--expiration-minutes", + dest="expiration_minutes", + type=int, + default=None, + help="Expiration time in minutes (-1 for no expiration).", + ) + p.add_argument( + "--auth-check", + dest="auth_check", + action="store_true", + default=None, + help="Require a ModelScope token when connecting to the remote URL.", + ) + p.add_argument( + "--env", + dest="env", + action="append", + default=None, + metavar="KEY=VALUE", + help="Environment variable for the MCP server (repeatable). " + "Available keys are listed in the server's env_schema (see `mcp info`).", ) add_subcmd_token_endpoint(p) p.set_defaults(_command=McpCommand, _mcp_leaf=_McpDeploy) @@ -102,13 +123,34 @@ def execute(self) -> None: payload["transport_type"] = self.args.transport_type if self.args.expiration_minutes is not None: payload["expiration_minutes"] = self.args.expiration_minutes - api.deploy_mcp_server(self.args.server_id, payload=payload or None) + if self.args.auth_check is not None: + payload["auth_check"] = self.args.auth_check + if self.args.env: + payload["env_info"] = parse_kv_pairs(self.args.env) + result = api.deploy_mcp_server(self.args.server_id, payload=payload or None) success(f"Deploy requested for MCP server: {self.args.server_id}") + _print_operational_url(result) + + +def _print_operational_url(result: object) -> None: + """Surface the deployed endpoint (McpOperationalUrl) when present.""" + if not isinstance(result, dict): + return + url = result.get("url") + if not url: + return + info(f" url: {url}") + if result.get("transport_type"): + info(f" transport: {result['transport_type']}") + if result.get("expiration"): + info(f" expires: {result['expiration']}") + if result.get("auth_required") is not None: + info(f" auth required: {result['auth_required']}") class _McpUndeploy(CLICommand): @staticmethod - def register(subparsers: Action) -> None: + def register(subparsers: SubParsers) -> None: p = subparsers.add_parser("undeploy", help="Undeploy an MCP server.") p.add_argument("server_id") add_subcmd_token_endpoint(p) diff --git a/src/modelscope_hub/cli/repo.py b/src/modelscope_hub/cli/repo.py index 88ae14a..c6ac7da 100644 --- a/src/modelscope_hub/cli/repo.py +++ b/src/modelscope_hub/cli/repo.py @@ -7,13 +7,22 @@ from __future__ import annotations import argparse -from argparse import Action from pathlib import Path from ..constants import RepoType from ..errors import AlreadyExistsError, is_repo_exists_error from ..types import RepoInfo -from .base import CLICommand, add_repo_type_arg, error, info, make_api, print_env_table, render_table, success +from .base import ( + CLICommand, + SubParsers, + add_repo_type_arg, + error, + info, + make_api, + print_env_table, + render_table, + success, +) from .compat import add_subcmd_token_endpoint @@ -44,7 +53,7 @@ class CreateCommand(CLICommand): """``ms create`` — create a new repository.""" @staticmethod - def register(subparsers: Action) -> None: + def register(subparsers: SubParsers) -> None: p = subparsers.add_parser("create", help="Create a new repository.") CreateCommand._add_arguments(p) p.set_defaults(_command=CreateCommand) @@ -57,16 +66,26 @@ def _add_arguments(p) -> None: p.add_argument("--license", dest="license", default=None) p.add_argument("--chinese-name", "--chinese_name", dest="chinese_name", default=None) p.add_argument("--description", dest="description", default=None) - p.add_argument("--exist-ok", "--exist_ok", dest="exist_ok", - action="store_true", default=False, - help="Do not error if repository already exists.") + p.add_argument( + "--exist-ok", + "--exist_ok", + dest="exist_ok", + action="store_true", + default=False, + help="Do not error if repository already exists.", + ) gated_group = p.add_mutually_exclusive_group() gated_group.add_argument( - "--gated", dest="gated", action="store_true", default=None, + "--gated", + dest="gated", + action="store_true", + default=None, help="Create a gated (application-required) repo. Implies private visibility.", ) gated_group.add_argument( - "--no-gated", dest="gated", action="store_false", + "--no-gated", + dest="gated", + action="store_false", help="Explicitly create a non-gated repo (default).", ) p.add_argument( @@ -81,16 +100,20 @@ def _add_arguments(p) -> None: p.add_argument("--cover-image", dest="cover_image", default=None, help="Studio cover image URL.") p.add_argument("--hardware", dest="hardware", default=None, help="Studio hardware spec.") p.add_argument( - "--category", dest="category", default=None, + "--category", + dest="category", + default=None, help="Skill category (required for skill repos). Options: " - "skill-management, developer-tools, marketing-seo, " - "frontend-development, ai-media, code-quality-testing, " - "mobile-development, cloud-devops, other.", + "skill-management, developer-tools, marketing-seo, " + "frontend-development, ai-media, code-quality-testing, " + "mobile-development, cloud-devops, other.", ) p.add_argument( - "--skill-file", dest="skill_file", default=None, + "--skill-file", + dest="skill_file", + default=None, help="Local zip for skill (max 5 MB, root must contain exactly one " - "SKILL.md with YAML front-matter: name, version, description).", + "SKILL.md with YAML front-matter: name, version, description).", ) add_subcmd_token_endpoint(p) @@ -143,7 +166,7 @@ class InfoCommand(CLICommand): """``ms info`` — show metadata for a repository.""" @staticmethod - def register(subparsers: Action) -> None: + def register(subparsers: SubParsers) -> None: p = subparsers.add_parser("info", help="Show metadata for a repository.") InfoCommand._add_arguments(p) p.set_defaults(_command=InfoCommand) @@ -164,7 +187,7 @@ class DeleteCommand(CLICommand): """``ms delete`` — delete a repository.""" @staticmethod - def register(subparsers: Action) -> None: + def register(subparsers: SubParsers) -> None: p = subparsers.add_parser("delete", help="Delete a repository (model or dataset).") DeleteCommand._add_arguments(p) p.set_defaults(_command=DeleteCommand) @@ -178,9 +201,11 @@ def _add_arguments(p) -> None: def execute(self) -> None: if not self.args.yes: - answer = input( - f"Delete {self.args.repo_type} {self.args.repo_id!r}? This cannot be undone. [y/N] " - ).strip().lower() + answer = ( + input(f"Delete {self.args.repo_type} {self.args.repo_id!r}? This cannot be undone. [y/N] ") + .strip() + .lower() + ) if answer not in ("y", "yes"): info("Aborted.") return @@ -193,7 +218,7 @@ class ListCommand(CLICommand): """``ms list`` — list repositories or environment variables.""" @staticmethod - def register(subparsers: Action) -> None: + def register(subparsers: SubParsers) -> None: p = subparsers.add_parser("list", help="List repositories or show configurable env vars.") ListCommand._add_arguments(p) p.set_defaults(_command=ListCommand) @@ -201,7 +226,9 @@ def register(subparsers: Action) -> None: @staticmethod def _add_arguments(p) -> None: p.add_argument( - "--envs", action="store_true", default=False, + "--envs", + action="store_true", + default=False, help="Show all configurable environment variables and exit.", ) add_repo_type_arg( @@ -218,8 +245,9 @@ def _add_arguments(p) -> None: p.add_argument("--owner", default=None) p.add_argument("--search", default=None, help=argparse.SUPPRESS) paging = p.add_mutually_exclusive_group() - paging.add_argument("--all", dest="fetch_all", action="store_true", default=False, - help="Fetch all pages automatically.") + paging.add_argument( + "--all", dest="fetch_all", action="store_true", default=False, help="Fetch all pages automatically." + ) paging.add_argument("--page", dest="page_number", type=int, default=1) p.add_argument("--page-size", dest="page_size", type=int, default=10) add_subcmd_token_endpoint(p) @@ -256,10 +284,7 @@ def execute(self) -> None: info("(no repositories found)") return self._render_table(result.items) - info( - f"\npage {result.page_number} / total {result.total_count} " - f"(page_size={result.page_size})" - ) + info(f"\npage {result.page_number} / total {result.total_count} (page_size={result.page_size})") def _fetch_all_pages(self, api) -> list[RepoInfo]: page_size = min(self.args.page_size, self._MAX_PAGE_SIZE) @@ -304,15 +329,14 @@ class RepoCommand(CLICommand): """Hidden compat dispatcher for ``ms repo create/info/list/delete``.""" @staticmethod - def register(subparsers: Action) -> None: + def register(subparsers: SubParsers) -> None: parser = subparsers.add_parser("repo") - try: - subparsers._choices_actions = [ - a for a in subparsers._choices_actions if a.dest != "repo" - ] - except AttributeError: - pass + # Hide the compat group from ``--help`` output. ``_choices_actions`` + # is private argparse state, so access it defensively. + choices_actions = getattr(subparsers, "_choices_actions", None) + if choices_actions is not None: + subparsers._choices_actions = [a for a in choices_actions if a.dest != "repo"] # type: ignore[attr-defined] sub = parser.add_subparsers(dest="repo_action", metavar="ACTION") sub.required = True diff --git a/src/modelscope_hub/cli/secret.py b/src/modelscope_hub/cli/secret.py index 3d1b275..5a09380 100644 --- a/src/modelscope_hub/cli/secret.py +++ b/src/modelscope_hub/cli/secret.py @@ -2,10 +2,8 @@ from __future__ import annotations -from argparse import Action - from ..constants import RepoType -from .base import CLICommand, add_repo_type_arg, info, make_api, render_table, success +from .base import CLICommand, SubParsers, add_repo_type_arg, info, make_api, render_table, success from .compat import add_subcmd_token_endpoint @@ -13,7 +11,7 @@ class SecretCommand(CLICommand): """Top-level dispatcher for the ``secret`` subcommands.""" @staticmethod - def register(subparsers: Action) -> None: + def register(subparsers: SubParsers) -> None: parser = subparsers.add_parser("secret", help="Manage repository secrets (studio).") sub = parser.add_subparsers(dest="secret_action", metavar="ACTION") sub.required = True @@ -43,7 +41,7 @@ def _add_studio_repo_type(parser) -> None: class _SecretList(CLICommand): @staticmethod - def register(subparsers: Action) -> None: + def register(subparsers: SubParsers) -> None: p = subparsers.add_parser("list", help="List secrets of a studio space.") p.add_argument("repo_id") _add_studio_repo_type(p) @@ -69,7 +67,7 @@ def execute(self) -> None: class _SecretAdd(CLICommand): @staticmethod - def register(subparsers: Action) -> None: + def register(subparsers: SubParsers) -> None: p = subparsers.add_parser("add", help="Add a new secret.") p.add_argument("repo_id") p.add_argument("key") @@ -86,7 +84,7 @@ def execute(self) -> None: class _SecretUpdate(CLICommand): @staticmethod - def register(subparsers: Action) -> None: + def register(subparsers: SubParsers) -> None: p = subparsers.add_parser("update", help="Update an existing secret.") p.add_argument("repo_id") p.add_argument("key") @@ -103,7 +101,7 @@ def execute(self) -> None: class _SecretDelete(CLICommand): @staticmethod - def register(subparsers: Action) -> None: + def register(subparsers: SubParsers) -> None: p = subparsers.add_parser("delete", help="Delete a secret.") p.add_argument("repo_id") p.add_argument("key") @@ -114,9 +112,7 @@ def register(subparsers: Action) -> None: def execute(self) -> None: if not self.args.yes: - answer = input( - f"Delete secret {self.args.key!r} from {self.args.repo_id}? [y/N] " - ).strip().lower() + answer = input(f"Delete secret {self.args.key!r} from {self.args.repo_id}? [y/N] ").strip().lower() if answer not in ("y", "yes"): info("Aborted.") return diff --git a/src/modelscope_hub/cli/upload.py b/src/modelscope_hub/cli/upload.py index cfe1714..3486eed 100644 --- a/src/modelscope_hub/cli/upload.py +++ b/src/modelscope_hub/cli/upload.py @@ -3,11 +3,10 @@ from __future__ import annotations import os -from argparse import Action from pathlib import Path from ..constants import RepoType -from .base import CLICommand, add_repo_type_arg, error, info, make_api, success +from .base import CLICommand, SubParsers, add_repo_type_arg, error, info, make_api, success from .compat import PatternAction, add_subcmd_token_endpoint @@ -19,7 +18,7 @@ class UploadCommand(CLICommand): """ @staticmethod - def register(subparsers: Action) -> None: + def register(subparsers: SubParsers) -> None: p = subparsers.add_parser( "upload", help="Upload a file or folder to a repository.", @@ -44,8 +43,12 @@ def register(subparsers: Action) -> None: required=False, ) p.add_argument("--commit-message", dest="commit_message", default=None) - p.add_argument("--commit-description", dest="commit_description", default=None, - help="Description for the generated commit.") + p.add_argument( + "--commit-description", + dest="commit_description", + default=None, + help="Description for the generated commit.", + ) p.add_argument("--revision", default=None, help="Target branch (default: master).") p.add_argument( "--include", diff --git a/src/modelscope_hub/compat/__init__.py b/src/modelscope_hub/compat/__init__.py index bee8a5c..64e34f3 100644 --- a/src/modelscope_hub/compat/__init__.py +++ b/src/modelscope_hub/compat/__init__.py @@ -22,14 +22,14 @@ FILE_HASH, MODELSCOPE_DOMAIN, MODELSCOPE_PREFER_AI_SITE, - ModelVisibility_INTERNAL, - ModelVisibility_PRIVATE, - ModelVisibility_PUBLIC, REPO_TYPE_DATASET, REPO_TYPE_MODEL, REPO_TYPE_STUDIO, REPO_TYPE_SUPPORT, TEMPORARY_FOLDER_NAME, + ModelVisibility_INTERNAL, + ModelVisibility_PRIVATE, + ModelVisibility_PUBLIC, ) from .file_download import dataset_file_download, model_file_download from .hub_api import LegacyHubApi diff --git a/src/modelscope_hub/compat/file_download.py b/src/modelscope_hub/compat/file_download.py index 7a098ee..58b1371 100644 --- a/src/modelscope_hub/compat/file_download.py +++ b/src/modelscope_hub/compat/file_download.py @@ -20,7 +20,7 @@ def _resolve_legacy_paths( repo_id: str, cache_dir: str | None, local_dir: str | None, - api: "HubApi", + api: HubApi, ) -> tuple[str | None, str | None]: """Resolve cache_dir/local_dir for legacy API compatibility. @@ -63,7 +63,10 @@ def model_file_download( except Exception: pass effective_cache, effective_local = _resolve_legacy_paths( - model_id, cache_dir, local_dir, api, + model_id, + cache_dir, + local_dir, + api, ) try: result = api.download_file( @@ -77,9 +80,7 @@ def model_file_download( user_agent=user_agent, ) except (NotExistError, AuthenticationError, PermissionDeniedError) as e: - raise _requests.exceptions.HTTPError( - str(e), response=getattr(e, 'response', None) - ) from e + raise _requests.exceptions.HTTPError(str(e), response=getattr(e, "response", None)) from e return str(result) @@ -112,7 +113,10 @@ def dataset_file_download( except Exception: pass effective_cache, effective_local = _resolve_legacy_paths( - dataset_id, cache_dir, local_dir, api, + dataset_id, + cache_dir, + local_dir, + api, ) try: result = api.download_file( @@ -126,7 +130,5 @@ def dataset_file_download( user_agent=user_agent, ) except (NotExistError, AuthenticationError, PermissionDeniedError) as e: - raise _requests.exceptions.HTTPError( - str(e), response=getattr(e, 'response', None) - ) from e + raise _requests.exceptions.HTTPError(str(e), response=getattr(e, "response", None)) from e return str(result) diff --git a/src/modelscope_hub/compat/hub_api.py b/src/modelscope_hub/compat/hub_api.py index 2fbb44b..9750d48 100644 --- a/src/modelscope_hub/compat/hub_api.py +++ b/src/modelscope_hub/compat/hub_api.py @@ -7,13 +7,12 @@ from __future__ import annotations import os +import time import warnings from collections import defaultdict from typing import TYPE_CHECKING, Any from urllib.parse import urlencode -import time - from ..api import HubApi from ..constants import RepoType from ..errors import ( @@ -33,7 +32,7 @@ DEFAULT_DATASET_REVISION = "master" -META_FILES_FORMAT = {'.json', '.csv', '.jsonl', '.tsv', '.py'} +META_FILES_FORMAT = {".json", ".csv", ".jsonl", ".tsv", ".py"} class LegacyHubApi: @@ -113,15 +112,15 @@ def get_model_files( handled by the configured token/session. """ files = self._api.list_repo_files( - model_id, RepoType.MODEL, revision=revision, recursive=recursive, + model_id, + RepoType.MODEL, + revision=revision, + recursive=recursive, ) result = [{"Path": f.path, "Size": f.size} for f in files] if root: prefix = root.strip("/") - result = [ - f for f in result - if f["Path"] == prefix or f["Path"].startswith(prefix + "/") - ] + result = [f for f in result if f["Path"] == prefix or str(f["Path"]).startswith(prefix + "/")] return result def create_repo( @@ -137,7 +136,7 @@ def create_repo( create_default_config: bool = False, endpoint: str | None = None, **kwargs: Any, - ) -> "RepoInfo | None": + ) -> RepoInfo | None: """Create a repository (legacy signature).""" api = self._api if token or endpoint: @@ -176,9 +175,7 @@ def create_model(self, model_id: str, **kwargs: Any) -> str: self.create_repo(model_id, repo_type="model", **kwargs) except (AuthenticationError, InvalidParameter) as e: if _is_auth_related(e): - raise ValueError( - "Token does not exist, please login first." - ) from e + raise ValueError("Token does not exist, please login first.") from e raise ep = self._endpoint or self._api._config.endpoint return f"{ep}/models/{model_id}" @@ -187,9 +184,7 @@ def push_model(self, model_id: str, model_dir: str, **kwargs: Any) -> None: """Upload a model directory (legacy signature).""" # Pre-validate model_dir if not os.path.isdir(model_dir): - raise ValueError( - f"model_dir '{model_dir}' does not exist or is not a directory." - ) + raise ValueError(f"model_dir '{model_dir}' does not exist or is not a directory.") config_files = ("configuration.json", "configuration.yaml", "configuration.yml") if not any(os.path.isfile(os.path.join(model_dir, f)) for f in config_files): logger.warning( @@ -234,9 +229,7 @@ def push_model(self, model_id: str, model_dir: str, **kwargs: Any) -> None: ) except (AuthenticationError, InvalidParameter) as e: if _is_auth_related(e): - raise ValueError( - "Token does not exist, please login first." - ) from e + raise ValueError("Token does not exist, please login first.") from e raise # ------------------------------------------------------------------ @@ -256,7 +249,9 @@ def get_endpoint_for_read( and ``MODELSCOPE_PREFER_AI_SITE`` env vars. """ return self._api.resolve_endpoint_for_read( - repo_id, repo_type=repo_type or "model", token=token, + repo_id, + repo_type=repo_type or "model", + token=token, ) def repo_exists( @@ -292,7 +287,7 @@ def list_repos( page_number: int = 1, page_size: int = 10, **filters: Any, - ) -> "PagedResult[RepoInfo]": + ) -> PagedResult[RepoInfo]: """List repositories of the given type. Delegates to :meth:`HubApi.list_repos`. @@ -313,7 +308,7 @@ def get_repo( repo_type: str | RepoType, *, revision: str | None = None, - ) -> "RepoInfo": + ) -> RepoInfo: """Get repository information. Delegates to :meth:`HubApi.get_repo`. @@ -342,9 +337,7 @@ def download_model( local_dir=local_dir, ) except (NotExistError, AuthenticationError, PermissionDeniedError) as e: - raise _requests.exceptions.HTTPError( - str(e), response=getattr(e, 'response', None) - ) from e + raise _requests.exceptions.HTTPError(str(e), response=getattr(e, "response", None)) from e return str(result) # ------------------------------------------------------------------ @@ -352,7 +345,9 @@ def download_model( # ------------------------------------------------------------------ def deploy_studio(self, studio_id: str, **kwargs: Any) -> dict: return self._api.deploy_repo( - studio_id, RepoType.STUDIO, payload=kwargs.get("payload"), + studio_id, + RepoType.STUDIO, + payload=kwargs.get("payload"), ) def stop_studio(self, studio_id: str, **kwargs: Any) -> dict: @@ -380,7 +375,9 @@ def delete_studio_secret(self, studio_id: str, key: str, **kwargs: Any) -> None: # Revision resolution # ------------------------------------------------------------------ def get_model_branches_and_tags_details( - self, model_id: str, **kwargs: Any, + self, + model_id: str, + **kwargs: Any, ) -> tuple[list[dict], list[dict]]: """Get model branches and tags as two separate detail lists. @@ -390,7 +387,9 @@ def get_model_branches_and_tags_details( return self._api.legacy.list_revisions_detail(model_id, "model") def get_model_branches_and_tags( - self, model_id: str, **kwargs: Any, + self, + model_id: str, + **kwargs: Any, ) -> tuple[list[str], list[str]]: """Get model branch and tag names.""" branches_detail, tags_detail = self.get_model_branches_and_tags_details(model_id) @@ -453,9 +452,7 @@ def _created_at(tag: dict) -> int: if revision is None: revision = "master" if revision not in all_branches and revision not in all_tags: - raise NotExistError( - f"The model: {model_id} has no revision: {revision}" - ) + raise NotExistError(f"The model: {model_id} has no revision: {revision}") detail = _find(tags_detail, revision) or _find(branches_detail, revision) return detail or {"Revision": revision} @@ -468,16 +465,11 @@ def _created_at(tag: dict) -> int: if not tags_detail: if revision is None or revision == "master": return _find(branches_detail, "master") or {"Revision": "master"} - raise NotExistError( - f"The model: {model_id} has no revision: {revision}" - ) + raise NotExistError(f"The model: {model_id} has no revision: {revision}") # Has tags if revision is None: - candidates = [ - t for t in tags_detail - if _created_at(t) <= release_timestamp - ] + candidates = [t for t in tags_detail if _created_at(t) <= release_timestamp] if candidates: return max(candidates, key=_created_at) return _find(branches_detail, "master") or {"Revision": "master"} @@ -488,10 +480,7 @@ def _created_at(tag: dict) -> int: if revision == "master": return _find(branches_detail, "master") or {"Revision": "master"} valid = ", ".join(all_tags) - raise NotExistError( - f"The model: {model_id} has no revision: {revision} " - f"(valid tags: {valid})" - ) + raise NotExistError(f"The model: {model_id} has no revision: {revision} (valid tags: {valid})") def get_valid_revision( self, @@ -502,7 +491,10 @@ def get_valid_revision( ) -> str: """Resolve a model revision to a concrete revision string.""" return self.get_valid_revision_detail( - model_id, revision=revision, cookies=cookies, endpoint=endpoint, + model_id, + revision=revision, + cookies=cookies, + endpoint=endpoint, )["Revision"] # ------------------------------------------------------------------ @@ -560,7 +552,7 @@ def list_datasets( search: str | None = None, endpoint: str | None = None, token: str | None = None, - ) -> "PagedResult": + ) -> PagedResult: """List datasets owned by a user/org. .. deprecated:: @@ -637,7 +629,7 @@ def get_dataset( *, endpoint: str | None = None, token: str | None = None, - ) -> "RepoInfo": + ) -> RepoInfo: """Get dataset information via OpenAPI. .. deprecated:: @@ -718,7 +710,8 @@ def get_dataset_files( else: raise ValueError(f"Invalid repo_id: {repo_id}") dataset_hub_id, _ = self.get_dataset_id_and_type( - dataset_name=_name, namespace=_owner, endpoint=endpoint, token=token) + dataset_name=_name, namespace=_owner, endpoint=endpoint, token=token + ) params: dict[str, Any] = { "Revision": revision, @@ -727,8 +720,7 @@ def get_dataset_files( "PageNumber": page_number, "PageSize": page_size, } - resp = api.legacy._request( - "GET", f"datasets/{dataset_hub_id}/repo/tree", params=params) + resp = api.legacy._request("GET", f"datasets/{dataset_hub_id}/repo/tree", params=params) data = api.legacy._json_data(resp) if isinstance(data, dict): return data.get("Files") or [] @@ -777,27 +769,28 @@ def get_dataset_meta_file_list( api = HubApi(endpoint=endpoint or self._endpoint, token=token) params = {"Revision": revision} - resp = api.legacy._request( - "GET", f"datasets/{dataset_id}/repo/tree", params=params) + resp = api.legacy._request("GET", f"datasets/{dataset_id}/repo/tree", params=params) data = api.legacy._json_data(resp) if data is None: raise NotExistError( f"The modelscope dataset [dataset_name = {dataset_name}, " - f"namespace = {namespace}, version = {revision}] does not exist") + f"namespace = {namespace}, version = {revision}] does not exist" + ) file_list = data.get("Files") if isinstance(data, dict) else data if file_list is None: raise NotExistError( f"The modelscope dataset [dataset_name = {dataset_name}, " - f"namespace = {namespace}, version = {revision}] does not exist") + f"namespace = {namespace}, version = {revision}] does not exist" + ) return file_list @staticmethod def dump_datatype_file(dataset_type: int, meta_cache_dir: str) -> None: """Dump dataset type marker file for offline formation detection.""" from modelscope.utils.constant import DatasetFormations + ext = DatasetFormations.formation_mark_ext.value - dataset_type_file_path = os.path.join( - meta_cache_dir, f"{str(dataset_type)}{ext}") + dataset_type_file_path = os.path.join(meta_cache_dir, f"{str(dataset_type)}{ext}") with open(dataset_type_file_path, "w") as fp: fp.write("*** Automatically-generated file, do not modify ***") @@ -872,12 +865,14 @@ def get_dataset_file_url( if not file_name or not dataset_name or not namespace: raise ValueError("Args (file_name, dataset_name, namespace) cannot be empty!") ep = endpoint or self._endpoint or self._api._config.endpoint - params = urlencode({ - "Source": "SDK", - "Revision": revision, - "FilePath": file_name, - "View": view, - }) + params = urlencode( + { + "Source": "SDK", + "Revision": revision, + "FilePath": file_name, + "View": view, + } + ) return f"{ep}/api/v1/datasets/{namespace}/{dataset_name}/repo?{params}" def get_dataset_file_url_origin( @@ -891,10 +886,7 @@ def get_dataset_file_url_origin( """Get dataset file URL, resolving meta files to API URLs.""" ep = endpoint or self._endpoint or self._api._config.endpoint if file_name and os.path.splitext(file_name)[-1] in META_FILES_FORMAT: - file_name = ( - f"{ep}/api/v1/datasets/{namespace}/{dataset_name}/repo?" - f"Revision={revision}&FilePath={file_name}" - ) + file_name = f"{ep}/api/v1/datasets/{namespace}/{dataset_name}/repo?Revision={revision}&FilePath={file_name}" return file_name def get_dataset_access_config( @@ -1011,6 +1003,7 @@ def _repo_info_to_dict(info: Any) -> dict: """Convert a RepoInfo to a plain dict with legacy PascalCase keys.""" if hasattr(info, "__dataclass_fields__"): from dataclasses import asdict + raw = asdict(info) elif hasattr(info, "__dict__"): raw = {k: v for k, v in info.__dict__.items() if not k.startswith("_")} diff --git a/src/modelscope_hub/compat/snapshot_download.py b/src/modelscope_hub/compat/snapshot_download.py index 00f84a5..7990508 100644 --- a/src/modelscope_hub/compat/snapshot_download.py +++ b/src/modelscope_hub/compat/snapshot_download.py @@ -8,7 +8,8 @@ from __future__ import annotations import warnings -from typing import TYPE_CHECKING, Sequence +from collections.abc import Sequence +from typing import TYPE_CHECKING import requests as _requests @@ -55,6 +56,7 @@ def snapshot_download( effective_id = repo_id or model_id if not effective_id: from ..errors import InvalidParameter + raise InvalidParameter("Please provide a valid model_id or repo_id") effective_type = repo_type or "model" @@ -72,13 +74,17 @@ def snapshot_download( if endpoint is None and not local_files_only: try: endpoint = api.resolve_endpoint_for_read( - effective_id, repo_type=effective_type, + effective_id, + repo_type=effective_type, ) api = HubApi(token=token, endpoint=endpoint) except Exception: pass effective_cache, effective_local = _resolve_legacy_paths( - effective_id, cache_dir, local_dir, api, + effective_id, + cache_dir, + local_dir, + api, ) try: result = api.download_repo( @@ -95,9 +101,7 @@ def snapshot_download( progress_callbacks=progress_callbacks, ) except (NotExistError, AuthenticationError, PermissionDeniedError) as e: - raise _requests.exceptions.HTTPError( - str(e), response=getattr(e, 'response', None) - ) from e + raise _requests.exceptions.HTTPError(str(e), response=getattr(e, "response", None)) from e return str(result) @@ -137,7 +141,10 @@ def dataset_snapshot_download( except Exception: pass effective_cache, effective_local = _resolve_legacy_paths( - dataset_id, cache_dir, local_dir, api, + dataset_id, + cache_dir, + local_dir, + api, ) try: result = api.download_repo( @@ -153,9 +160,7 @@ def dataset_snapshot_download( user_agent=user_agent, ) except (NotExistError, AuthenticationError, PermissionDeniedError) as e: - raise _requests.exceptions.HTTPError( - str(e), response=getattr(e, 'response', None) - ) from e + raise _requests.exceptions.HTTPError(str(e), response=getattr(e, "response", None)) from e return str(result) diff --git a/src/modelscope_hub/config.py b/src/modelscope_hub/config.py index 5067011..4d3056f 100644 --- a/src/modelscope_hub/config.py +++ b/src/modelscope_hub/config.py @@ -17,6 +17,7 @@ import warnings from dataclasses import dataclass, field from pathlib import Path +from typing import Any from .constants import ( CONFIG_DIR_NAME, @@ -52,15 +53,9 @@ class HubConfig: endpoint: str | None = None # type: ignore[assignment] # sentinel; always str after __post_init__ cache_dir: Path = field( - default_factory=lambda: _expand( - os.environ.get(ENV_CACHE) or Path.home() / ".cache" / DEFAULT_CACHE_DIR_NAME - ) - ) - config_dir: Path = field( - default_factory=lambda: _expand( - os.environ.get(ENV_HOME) or Path.home() / CONFIG_DIR_NAME - ) + default_factory=lambda: _expand(os.environ.get(ENV_CACHE) or Path.home() / ".cache" / DEFAULT_CACHE_DIR_NAME) ) + config_dir: Path = field(default_factory=lambda: _expand(os.environ.get(ENV_HOME) or Path.home() / CONFIG_DIR_NAME)) token: str | None = None _logged_out: bool = field(default=False, init=False, repr=False) _endpoint_overridden: bool = field(default=False, init=False, repr=False) @@ -80,8 +75,7 @@ def __post_init__(self) -> None: domain = os.environ.get(ENV_MODELSCOPE_DOMAIN, "").strip() if domain: warnings.warn( - "Environment variable MODELSCOPE_DOMAIN is deprecated, " - "use MODELSCOPE_ENDPOINT instead.", + "Environment variable MODELSCOPE_DOMAIN is deprecated, use MODELSCOPE_ENDPOINT instead.", FutureWarning, stacklevel=2, ) @@ -94,7 +88,7 @@ def __post_init__(self) -> None: # Ensure endpoint always has a scheme if self.endpoint and not self.endpoint.startswith(("http://", "https://")): self.endpoint = f"https://{self.endpoint}" - self.endpoint = self.endpoint.rstrip("/") + self.endpoint = (self.endpoint or DEFAULT_ENDPOINT).rstrip("/") # 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,22 +132,36 @@ def save_token(self, token: str) -> None: import time from http.cookiejar import Cookie - from requests.cookies import RequestsCookieJar from urllib.parse import urlparse + from requests.cookies import RequestsCookieJar + token = token.strip() - domain = urlparse(self.endpoint).hostname or "modelscope.cn" + domain = urlparse(self.endpoint or DEFAULT_ENDPOINT).hostname or "modelscope.cn" expires = int(time.time()) + 30 * 24 * 3600 # 30 days jar = RequestsCookieJar() - jar.set_cookie(Cookie( - version=0, name="m_session_id", value=token, - port=None, port_specified=False, - domain=domain, domain_specified=True, domain_initial_dot=False, - path="/", path_specified=True, - secure=False, expires=expires, discard=False, - comment=None, comment_url=None, rest={}, rfc2109=False, - )) + jar.set_cookie( + Cookie( + version=0, + name="m_session_id", + value=token, + port=None, + port_specified=False, + domain=domain, + domain_specified=True, + domain_initial_dot=False, + path="/", + path_specified=True, + secure=False, + expires=expires, + discard=False, + comment=None, + comment_url=None, + rest={}, + rfc2109=False, + ) + ) self.save_cookies(jar) self.token = token self._logged_out = False @@ -201,7 +209,7 @@ def save_cookies(self, cookies: object) -> None: pickle.dump(cookies, f) path.chmod(stat.S_IRUSR | stat.S_IWUSR) - def load_cookies(self) -> object | None: + def load_cookies(self) -> Any: """Load saved cookies, returning None if absent or expired.""" import pickle diff --git a/src/modelscope_hub/constants.py b/src/modelscope_hub/constants.py index 5b7e02f..aec9061 100644 --- a/src/modelscope_hub/constants.py +++ b/src/modelscope_hub/constants.py @@ -8,17 +8,20 @@ from __future__ import annotations import os +import sys from dataclasses import dataclass from enum import Enum, IntEnum - # --------------------------------------------------------------------------- # StrEnum compatibility shim (Python 3.10 lacks :class:`enum.StrEnum`). +# ``sys.version_info`` branching (instead of try/except) lets type checkers +# resolve the correct definition statically. # --------------------------------------------------------------------------- -try: # pragma: no cover - exercised implicitly by the import path - from enum import StrEnum # type: ignore[attr-defined] -except ImportError: # Python 3.10 - class StrEnum(str, Enum): # type: ignore[no-redef] +if sys.version_info >= (3, 11): + from enum import StrEnum +else: + + class StrEnum(str, Enum): """Minimal backport of :class:`enum.StrEnum` for Python 3.10.""" def __str__(self) -> str: # noqa: D401 - mirror stdlib behaviour @@ -38,10 +41,16 @@ class EnvVar: category: str # Core, Network, Download, Upload, Logging, Deprecated deprecated_names: tuple[str, ...] = () + ENV_REGISTRY: list[EnvVar] = [] CATEGORY_ORDER: tuple[str, ...] = ( - "Core", "Network", "Download", "Upload", "Logging", "Deprecated", + "Core", + "Network", + "Download", + "Upload", + "Logging", + "Deprecated", ) @@ -75,7 +84,7 @@ def label(self) -> str: return self.name.lower() @classmethod - def from_label(cls, label: str) -> "Visibility": + def from_label(cls, label: str) -> Visibility: """Resolve a visibility from its lowercase label or numeric string. Supports both label strings ('private', 'internal', 'public') and @@ -139,9 +148,9 @@ def _env(name: str, *deprecated_names: str) -> str | None: value = os.environ.get(old) if value is not None: import warnings + warnings.warn( - f"Environment variable {old!r} is deprecated, " - f"use {name!r} instead.", + f"Environment variable {old!r} is deprecated, use {name!r} instead.", FutureWarning, stacklevel=4, ) @@ -159,8 +168,7 @@ def _env_int( """Read a positive integer from the environment and register it.""" all_deprecated = deprecated_names or _DEPRECATED_LOOKUP.get(name, ()) if description and category: - _env_register(name, str(default), description, category, - deprecated_names=all_deprecated) + _env_register(name, str(default), description, category, deprecated_names=all_deprecated) raw = _env(name, *all_deprecated) if raw is None or raw.strip() == "": return default @@ -186,8 +194,7 @@ def _env_int_mb( """ all_deprecated = deprecated_byte_names or _DEPRECATED_LOOKUP.get(name, ()) if description and category: - _env_register(name, str(default_mb), description, category, - deprecated_names=all_deprecated) + _env_register(name, str(default_mb), description, category, deprecated_names=all_deprecated) # Check the new name first (value in MB) raw = os.environ.get(name) if raw is not None and raw.strip(): @@ -201,6 +208,7 @@ def _env_int_mb( raw = os.environ.get(old) if raw is not None and raw.strip(): import warnings + warnings.warn( f"Environment variable {old!r} is deprecated, " f"use {name!r} instead. Note: {name!r} expects a value in MB.", @@ -224,8 +232,7 @@ def _env_bool( ) -> bool: """Read a boolean from the environment and register it.""" if description and category: - _env_register(name, str(default).lower(), description, category, - deprecated_names=deprecated_names) + _env_register(name, str(default).lower(), description, category, deprecated_names=deprecated_names) all_deprecated = deprecated_names or _DEPRECATED_LOOKUP.get(name, ()) raw = _env(name, *all_deprecated) if raw is None or raw.strip() == "": @@ -266,19 +273,25 @@ def _env_register( # Network / IO tunables # --------------------------------------------------------------------------- API_TIMEOUT: int = _env_int( - "MODELSCOPE_API_TIMEOUT", 60, - "HTTP request timeout (seconds)", "Network", + "MODELSCOPE_API_TIMEOUT", + 60, + "HTTP request timeout (seconds)", + "Network", "API_TIMEOUT", ) API_CONNECT_TIMEOUT: int = _env_int( - "MODELSCOPE_API_CONNECT_TIMEOUT", 10, - "HTTP connect timeout (seconds)", "Network", + "MODELSCOPE_API_CONNECT_TIMEOUT", + 10, + "HTTP connect timeout (seconds)", + "Network", ) API_MAX_RETRIES: int = _env_int( - "MODELSCOPE_API_MAX_RETRIES", 5, - "Max retry attempts for transient failures", "Network", + "MODELSCOPE_API_MAX_RETRIES", + 5, + "Max retry attempts for transient failures", + "Network", "API_MAX_RETRIES", ) @@ -298,38 +311,54 @@ def _env_register( # Download tunables # --------------------------------------------------------------------------- DOWNLOAD_CHUNK_SIZE: int = _env_int_mb( - "MODELSCOPE_DOWNLOAD_CHUNK_SIZE_MB", 1, - "Streaming chunk size (MB)", "Download", + "MODELSCOPE_DOWNLOAD_CHUNK_SIZE_MB", + 1, + "Streaming chunk size (MB)", + "Download", "DOWNLOAD_CHUNK_SIZE", ) -DOWNLOAD_PARALLEL_THRESHOLD: int = _env_int( - "MODELSCOPE_DOWNLOAD_PARALLEL_THRESHOLD_MB", 500, - "Parallel download threshold (MB)", "Download", - "MODELSCOPE_PARALLEL_DOWNLOAD_THRESHOLD_MB", -) * 1024 * 1024 +DOWNLOAD_PARALLEL_THRESHOLD: int = ( + _env_int( + "MODELSCOPE_DOWNLOAD_PARALLEL_THRESHOLD_MB", + 500, + "Parallel download threshold (MB)", + "Download", + "MODELSCOPE_PARALLEL_DOWNLOAD_THRESHOLD_MB", + ) + * 1024 + * 1024 +) DOWNLOAD_PARALLELS: int = _env_int( - "MODELSCOPE_DOWNLOAD_PARALLEL_WORKERS", 1, - "Parallel range-download streams", "Download", + "MODELSCOPE_DOWNLOAD_PARALLEL_WORKERS", + 1, + "Parallel range-download streams", + "Download", "MODELSCOPE_DOWNLOAD_PARALLELS", ) DOWNLOAD_RETRY_TIMES: int = _env_int( - "MODELSCOPE_DOWNLOAD_MAX_RETRIES", 5, - "Per-file download retry count", "Download", + "MODELSCOPE_DOWNLOAD_MAX_RETRIES", + 5, + "Per-file download retry count", + "Download", "DOWNLOAD_RETRY_TIMES", ) DOWNLOAD_TIMEOUT: int = _env_int( - "MODELSCOPE_DOWNLOAD_TIMEOUT", 60, - "Per-file download timeout (seconds)", "Download", + "MODELSCOPE_DOWNLOAD_TIMEOUT", + 60, + "Per-file download timeout (seconds)", + "Download", "DOWNLOAD_TIMEOUT", ) DOWNLOAD_PART_SIZE: int = _env_int_mb( - "MODELSCOPE_DOWNLOAD_PART_SIZE_MB", 160, - "Parallel range chunk size (MB)", "Download", + "MODELSCOPE_DOWNLOAD_PART_SIZE_MB", + 160, + "Parallel range chunk size (MB)", + "Download", "DOWNLOAD_PART_SIZE", ) @@ -340,20 +369,36 @@ def _env_register( """API response field name for file hash.""" ENV_FILE_LOCK: str = "MODELSCOPE_DOWNLOAD_FILE_LOCK" -_env_register(ENV_FILE_LOCK, "true", "File lock for multiprocess download safety", "Download", - deprecated_names=("MODELSCOPE_HUB_FILE_LOCK",)) +_env_register( + ENV_FILE_LOCK, + "true", + "File lock for multiprocess download safety", + "Download", + deprecated_names=("MODELSCOPE_HUB_FILE_LOCK",), +) ENV_INTRA_CLOUD_ACCELERATION: str = "MODELSCOPE_DOWNLOAD_INTRA_CLOUD" -_env_register(ENV_INTRA_CLOUD_ACCELERATION, "true", "Alibaba cloud intra-cloud acceleration", "Download", - deprecated_names=("INTRA_CLOUD_ACCELERATION",)) +_env_register( + ENV_INTRA_CLOUD_ACCELERATION, + "true", + "Alibaba cloud intra-cloud acceleration", + "Download", + deprecated_names=("INTRA_CLOUD_ACCELERATION",), +) ENV_INTRA_CLOUD_REGION: str = "MODELSCOPE_DOWNLOAD_INTRA_CLOUD_REGION" -_env_register(ENV_INTRA_CLOUD_REGION, "(auto)", "Override intra-cloud region ID", "Download", - deprecated_names=("INTRA_CLOUD_ACCELERATION_REGION",)) +_env_register( + ENV_INTRA_CLOUD_REGION, + "(auto)", + "Override intra-cloud region ID", + "Download", + deprecated_names=("INTRA_CLOUD_ACCELERATION_REGION",), +) ENV_INTER_CLOUD_REGIONS: str = "MODELSCOPE_DOWNLOAD_INTER_CLOUD_REGIONS" -_env_register(ENV_INTER_CLOUD_REGIONS, "", - "Comma-separated peer regions for cross-region internal acceleration", "Download") +_env_register( + ENV_INTER_CLOUD_REGIONS, "", "Comma-separated peer regions for cross-region internal acceleration", "Download" +) UPLOAD_LFS_THRESHOLD: int = _env_int("UPLOAD_LFS_THRESHOLD", 5 * 1024 * 1024) UPLOAD_LFS_ENFORCE_THRESHOLD: int = _env_int("UPLOAD_LFS_ENFORCE_THRESHOLD", 1 * 1024 * 1024) @@ -366,21 +411,23 @@ def _env_register( # Upload: blob timeout UPLOAD_BLOB_CONNECT_TIMEOUT: int = _env_int( - "MODELSCOPE_UPLOAD_CONNECT_TIMEOUT", 30, - "Upload connect timeout (seconds)", "Upload", + "MODELSCOPE_UPLOAD_CONNECT_TIMEOUT", + 30, + "Upload connect timeout (seconds)", + "Upload", "UPLOAD_BLOB_CONNECT_TIMEOUT", ) UPLOAD_BLOB_READ_TIMEOUT: int = _env_int( - "MODELSCOPE_UPLOAD_READ_TIMEOUT", 3600, - "Upload read timeout (seconds)", "Upload", + "MODELSCOPE_UPLOAD_READ_TIMEOUT", + 3600, + "Upload read timeout (seconds)", + "Upload", "UPLOAD_BLOB_READ_TIMEOUT", ) # Upload: urllib3 retry UPLOAD_RETRY_ALLOWED_METHODS: frozenset[str] = frozenset( - os.environ.get( - "UPLOAD_RETRY_ALLOWED_METHODS", "GET,HEAD,DELETE,OPTIONS,TRACE" - ).split(",") + os.environ.get("UPLOAD_RETRY_ALLOWED_METHODS", "GET,HEAD,DELETE,OPTIONS,TRACE").split(",") ) # Upload: batching @@ -415,15 +462,19 @@ def _env_register( # Upload: workers DEFAULT_MAX_WORKERS: int = _env_int( - "MODELSCOPE_UPLOAD_MAX_WORKERS", min(8, (os.cpu_count() or 4) + 4), - "Default parallel worker threads (min(8, cpu+4))", "Upload", + "MODELSCOPE_UPLOAD_MAX_WORKERS", + min(8, (os.cpu_count() or 4) + 4), + "Default parallel worker threads (min(8, cpu+4))", + "Upload", "DEFAULT_MAX_WORKERS", ) # Upload: cache / tracker UPLOAD_USE_CACHE: bool = _env_bool( - "MODELSCOPE_UPLOAD_CACHE", True, - "Enable resumable upload cache", "Upload", + "MODELSCOPE_UPLOAD_CACHE", + True, + "Enable resumable upload cache", + "Upload", "UPLOAD_USE_CACHE", ) UPLOAD_CACHE_FILE: str = ".ms_upload_cache" @@ -431,13 +482,17 @@ def _env_register( # Upload: limits UPLOAD_MAX_FILE_SIZE: int = _env_int_mb( - "MODELSCOPE_UPLOAD_MAX_FILE_SIZE_MB", 100 * 1024, - "Max single file size (MB, default 100 GB)", "Upload", + "MODELSCOPE_UPLOAD_MAX_FILE_SIZE_MB", + 100 * 1024, + "Max single file size (MB, default 100 GB)", + "Upload", "UPLOAD_MAX_FILE_SIZE", ) UPLOAD_MAX_FILE_COUNT: int = _env_int( - "MODELSCOPE_UPLOAD_MAX_FILE_COUNT", 100_000, - "Max total files per upload", "Upload", + "MODELSCOPE_UPLOAD_MAX_FILE_COUNT", + 100_000, + "Max total files per upload", + "Upload", "UPLOAD_MAX_FILE_COUNT", ) UPLOAD_MAX_FILE_COUNT_IN_DIR: int = _env_int("UPLOAD_MAX_FILE_COUNT_IN_DIR", 50_000) @@ -445,24 +500,94 @@ def _env_register( # LFS suffix lists (from old SDK — determines upload mode regardless of size) MODEL_LFS_SUFFIX: list[str] = [ - ".7z", ".arrow", ".bin", ".bz2", ".ckpt", ".ftz", ".gz", ".h5", - ".joblib", ".mlmodel", ".model", ".msgpack", ".npy", ".npz", ".onnx", - ".ot", ".parquet", ".pb", ".pickle", ".pkl", ".pt", ".pth", ".rar", - ".safetensors", ".tar", ".tflite", ".tgz", ".wasm", ".xz", ".zip", ".zst", + ".7z", + ".arrow", + ".bin", + ".bz2", + ".ckpt", + ".ftz", + ".gz", + ".h5", + ".joblib", + ".mlmodel", + ".model", + ".msgpack", + ".npy", + ".npz", + ".onnx", + ".ot", + ".parquet", + ".pb", + ".pickle", + ".pkl", + ".pt", + ".pth", + ".rar", + ".safetensors", + ".tar", + ".tflite", + ".tgz", + ".wasm", + ".xz", + ".zip", + ".zst", ] DATASET_LFS_SUFFIX: list[str] = [ - ".7z", ".aac", ".arrow", ".audio", ".bmp", ".bin", ".bz2", ".flac", - ".ftz", ".gif", ".gz", ".h5", ".jack", ".jpeg", ".jpg", ".png", ".jsonl", - ".joblib", ".lz4", ".msgpack", ".npy", ".npz", ".ot", ".parquet", ".pb", - ".pickle", ".pcm", ".pkl", ".raw", ".rar", ".sam", ".tar", ".tgz", - ".wasm", ".wav", ".webm", ".webp", ".zip", ".zst", ".tiff", ".mp3", - ".mp4", ".ogg", + ".7z", + ".aac", + ".arrow", + ".audio", + ".bmp", + ".bin", + ".bz2", + ".flac", + ".ftz", + ".gif", + ".gz", + ".h5", + ".jack", + ".jpeg", + ".jpg", + ".png", + ".jsonl", + ".joblib", + ".lz4", + ".msgpack", + ".npy", + ".npz", + ".ot", + ".parquet", + ".pb", + ".pickle", + ".pcm", + ".pkl", + ".raw", + ".rar", + ".sam", + ".tar", + ".tgz", + ".wasm", + ".wav", + ".webm", + ".webp", + ".zip", + ".zst", + ".tiff", + ".mp3", + ".mp4", + ".ogg", ] # Default ignore patterns for folder upload DEFAULT_IGNORE_PATTERNS: list[str] = [ - ".git", ".git/*", "*/.git", "**/.git/**", - ".cache", ".cache/*", "*/.cache", "**/.cache/**", + ".git", + ".git/*", + "*/.git", + "**/.git/**", + ".cache", + ".cache/*", + "*/.cache", + "**/.cache/**", ] @@ -486,8 +611,13 @@ def _env_register( # Logging / deprecated (read logic in utils/logger.py, cli/compat.py) # --------------------------------------------------------------------------- _env_register("MODELSCOPE_LOG_LEVEL", "INFO", "SDK log level (DEBUG/INFO/WARNING/ERROR)", "Logging") -_env_register("MODELSCOPE_NO_DEPRECATION_WARNINGS", "-", "Suppress deprecation warnings", "Logging", - deprecated_names=("MODELSCOPE_HUB_NO_DEPRECATION_WARNINGS",)) +_env_register( + "MODELSCOPE_NO_DEPRECATION_WARNINGS", + "-", + "Suppress deprecation warnings", + "Logging", + deprecated_names=("MODELSCOPE_HUB_NO_DEPRECATION_WARNINGS",), +) # --------------------------------------------------------------------------- diff --git a/src/modelscope_hub/errors.py b/src/modelscope_hub/errors.py index 71e7fd3..b4c9db3 100644 --- a/src/modelscope_hub/errors.py +++ b/src/modelscope_hub/errors.py @@ -19,7 +19,7 @@ import re from typing import TYPE_CHECKING, Any -from urllib.parse import parse_qs, urlencode, urlparse, urlunparse +from urllib.parse import parse_qs, urlparse, urlunparse if TYPE_CHECKING: # pragma: no cover - type-only imports from requests import Response @@ -28,16 +28,36 @@ # Credential redaction helpers # --------------------------------------------------------------------------- _SENSITIVE_KEYWORDS: tuple[str, ...] = ( - "token", "secret", "password", "cookie", "authorization", - "credential", "session", "api_key", "apikey", + "token", + "secret", + "password", + "cookie", + "authorization", + "credential", + "session", + "api_key", + "apikey", +) +_SENSITIVE_QUERY_KEYS: frozenset[str] = frozenset( + { + "token", + "access_token", + "auth_token", + "api_key", + "apikey", + "cookie", + "m_session_id", + "session", + "secret", + "password", + "key", + "authorization", + "credentials", + } ) -_SENSITIVE_QUERY_KEYS: frozenset[str] = frozenset({ - "token", "access_token", "auth_token", "api_key", "apikey", - "cookie", "m_session_id", "session", - "secret", "password", "key", "authorization", "credentials", -}) _SENSITIVE_BODY_KEYS: re.Pattern[str] = re.compile( - "|".join(_SENSITIVE_KEYWORDS), re.IGNORECASE, + "|".join(_SENSITIVE_KEYWORDS), + re.IGNORECASE, ) _REDACTED = "***" @@ -64,10 +84,7 @@ def _redact_url(url: str) -> str: def _redact_body(body: Any) -> Any: """Deep-redact sensitive keys in a response body structure.""" if isinstance(body, dict): - return { - k: _REDACTED if _SENSITIVE_BODY_KEYS.search(k) else _redact_body(v) - for k, v in body.items() - } + return {k: _REDACTED if _SENSITIVE_BODY_KEYS.search(k) else _redact_body(v) for k, v in body.items()} if isinstance(body, list): return [_redact_body(item) for item in body] return body @@ -306,9 +323,7 @@ class StorageError(HubError): error_code = "E1003" retryable = True - suggestion = ( - "File upload/download failed (storage service error). Please retry later." - ) + suggestion = "File upload/download failed (storage service error). Please retry later." class FileIntegrityError(HubError): @@ -388,8 +403,10 @@ class NotSupportedError(HubError): "模型不存在": "Model does not exist.", "数据集不存在": "Dataset does not exist.", "创建空间失败": "Failed to create studio.", - "the current token no longer supports deletion operations. Please go to the site page : https://www.modelscope.cn to delete": - "Deletion is restricted to web console. Visit https://modelscope.cn to delete.", + "the current token no longer supports deletion operations. " + "Please go to the site page : https://www.modelscope.cn to delete": ( + "Deletion is restricted to web console. Visit https://modelscope.cn to delete." + ), } @@ -403,7 +420,7 @@ def _translate_message(msg: str) -> str: return msg -def _extract_payload(response: "Response") -> tuple[str, str | None, Any | None]: +def _extract_payload(response: Response) -> tuple[str, str | None, Any | None]: """Best-effort extraction of (message, request_id, body) from a response.""" request_id = response.headers.get("x-request-id") or response.headers.get("X-Request-Id") body: Any | None = None @@ -429,14 +446,11 @@ def _extract_payload(response: "Response") -> tuple[str, str | None, Any | None] if isinstance(value, str) and value.strip(): message = value.strip() break - request_id = ( - body.get("request_id") or body.get("requestId") - or body.get("RequestId") or request_id - ) + request_id = body.get("request_id") or body.get("requestId") or body.get("RequestId") or request_id return _translate_message(message), request_id, body -def raise_for_status(response: "Response") -> None: +def raise_for_status(response: Response) -> None: """Inspect ``response`` and raise the most specific exception on failure. Parameters @@ -468,8 +482,7 @@ def raise_for_status(response: "Response") -> None: # Detect "already exists" errors before falling back to InvalidParameter if exc_cls is InvalidParameter and isinstance(body, dict): code = body.get("Code") or body.get("code") - msg_text = (body.get("Message") or body.get("message") - or body.get("msg") or body.get("Msg") or "").lower() + msg_text = (body.get("Message") or body.get("message") or body.get("msg") or body.get("Msg") or "").lower() is_exists = False if code is not None: try: @@ -511,20 +524,22 @@ def raise_for_status(response: "Response") -> None: # Repo-exists detection (shared by cli/repo.py and compat/hub_api.py) # --------------------------------------------------------------------------- _ALREADY_EXISTS_CODES: set[int] = { - 10020101001, # 国内站 - 数据集已存在 - 10010101001, # 国内站 - 模型已存在 - 10010202004, # 国际站 - 名称已被使用 + 10020101001, # 国内站 - 数据集已存在 + 10010101001, # 国内站 - 模型已存在 + 10010202004, # 国际站 - 名称已被使用 } -_ALREADY_EXISTS_KEYWORDS: frozenset[str] = frozenset({ - "exist", - "already", - "can not be used", - "not available", - "已被注册", - "已存在", - "名称不可用", -}) +_ALREADY_EXISTS_KEYWORDS: frozenset[str] = frozenset( + { + "exist", + "already", + "can not be used", + "not available", + "已被注册", + "已存在", + "名称不可用", + } +) def is_repo_exists_error(exc: BaseException) -> bool: @@ -545,7 +560,7 @@ def is_repo_exists_error(exc: BaseException) -> bool: if isinstance(body, dict): code = body.get("Code") or body.get("code") try: - if int(code) in _ALREADY_EXISTS_CODES: + if code is not None and int(code) in _ALREADY_EXISTS_CODES: return True except (TypeError, ValueError): pass diff --git a/src/modelscope_hub/types.py b/src/modelscope_hub/types.py index d6db646..f690d49 100644 --- a/src/modelscope_hub/types.py +++ b/src/modelscope_hub/types.py @@ -7,10 +7,11 @@ from __future__ import annotations +from collections.abc import Mapping from dataclasses import dataclass, field, fields from datetime import datetime, timezone from enum import Enum -from typing import Any, Generic, Mapping, Type, TypedDict, TypeVar +from typing import Any, Generic, TypedDict, TypeVar from .constants import RepoType, Visibility @@ -40,7 +41,7 @@ class _FromDictMixin: """Adds tolerant ``from_dict`` construction to a dataclass.""" @classmethod - def from_dict(cls: Type[_TDataclass], data: Mapping[str, Any] | None) -> _TDataclass: + def from_dict(cls: type[_TDataclass], data: Mapping[str, Any] | None) -> _TDataclass: if not data: return cls() # type: ignore[call-arg] known = {f.name for f in fields(cls)} # type: ignore[arg-type] @@ -213,7 +214,7 @@ class CachedRepoInfo(_FromDictMixin): revision: str | None = None size_on_disk: int = 0 nb_files: int = 0 - last_accessed: datetime | str | int | None = None + last_accessed: datetime | str | int | float | None = None local_path: str | None = None def __post_init__(self) -> None: diff --git a/src/modelscope_hub/utils/file_utils.py b/src/modelscope_hub/utils/file_utils.py index 7c9e396..38285c8 100644 --- a/src/modelscope_hub/utils/file_utils.py +++ b/src/modelscope_hub/utils/file_utils.py @@ -6,13 +6,13 @@ import io import os from pathlib import Path -from typing import IO, Union +from typing import IO from ..constants import DEFAULT_CACHE_DIR_NAME, DOWNLOAD_CHUNK_SIZE, ENV_CACHE from ..errors import FileIntegrityError -PathLike = Union[str, os.PathLike[str], Path] -FileObj = Union[IO[bytes], io.IOBase] +PathLike = str | os.PathLike[str] | Path +FileObj = IO[bytes] | io.IOBase def compute_hash( diff --git a/src/modelscope_hub/utils/format.py b/src/modelscope_hub/utils/format.py index 2f3f0ac..4e0408e 100644 --- a/src/modelscope_hub/utils/format.py +++ b/src/modelscope_hub/utils/format.py @@ -3,7 +3,7 @@ from __future__ import annotations import time -from typing import Iterable, Sequence +from collections.abc import Iterable, Sequence # --------------------------------------------------------------------------- # Size formatting @@ -101,8 +101,7 @@ def tabulate( raise ValueError(f"max_width must be >= 1, got {max_width}") ncols = len(headers) str_rows: list[list[str]] = [ - [_cell(row[i] if i < len(row) else "", max_width) for i in range(ncols)] - for row in rows + [_cell(row[i] if i < len(row) else "", max_width) for i in range(ncols)] for row in rows ] widths = [len(h) for h in headers] diff --git a/src/modelscope_hub/utils/media.py b/src/modelscope_hub/utils/media.py index 8aedec2..038ca49 100644 --- a/src/modelscope_hub/utils/media.py +++ b/src/modelscope_hub/utils/media.py @@ -52,9 +52,7 @@ def encode_media_to_base64(media_file_path: str | os.PathLike) -> str: if mime_type is None: mime_type = _FALLBACK_MIME_TYPES.get(path.suffix.lower()) if mime_type is None: - raise ValueError( - f"Cannot determine MIME type for file: {path}" - ) + raise ValueError(f"Cannot determine MIME type for file: {path}") encoded = base64.b64encode(path.read_bytes()).decode("ascii") return f"data:{mime_type};base64,{encoded}" diff --git a/src/modelscope_hub/utils/time_utils.py b/src/modelscope_hub/utils/time_utils.py index 654380a..3393fef 100644 --- a/src/modelscope_hub/utils/time_utils.py +++ b/src/modelscope_hub/utils/time_utils.py @@ -5,7 +5,6 @@ import re import zoneinfo from datetime import datetime, timezone -from typing import Union # Accepted ISO-like formats for naive (local) timestamps. _NAIVE_FORMATS: tuple[str, ...] = ( @@ -56,7 +55,7 @@ def _parse_naive_string(value: str, target_tz: zoneinfo.ZoneInfo) -> datetime: def parse_timestamp( - value: Union[int, str, datetime, None], + value: int | str | datetime | None, *, tz: str = "Asia/Shanghai", ) -> datetime | None: diff --git a/tests/agent/test_agent_cli.py b/tests/agent/test_agent_cli.py index 70fc0fb..03b0c4a 100644 --- a/tests/agent/test_agent_cli.py +++ b/tests/agent/test_agent_cli.py @@ -6,6 +6,7 @@ framework-aware commands (convert/watch/status/backups/restore/stop) now live in modelscope-agent. """ + from __future__ import annotations import argparse @@ -62,8 +63,9 @@ def commit_files(self, path, name, actions, revision="master", commit_message="s type(self).commits.append(actions) return {"ok": True} - def upload_lfs_file(self, path, name, file_path, content, action="create", - revision="master", commit_message="sync"): + def upload_lfs_file( + self, path, name, file_path, content, action="create", revision="master", commit_message="sync" + ): type(self).lfs_uploads.append((file_path, content)) return {"ok": True} @@ -118,8 +120,13 @@ def test_list_empty(self): @mock.patch.object(cli_agent, "AgentApi", _StubClient) def test_list_rows(self): _StubClient.agents = [ - {"Path": "user", "Name": "a1", "Framework": "qoder", - "Visibility": "public", "LastUpdatedDate": "2024-01-02T03:04:05"}, + { + "Path": "user", + "Name": "a1", + "Framework": "qoder", + "Visibility": "public", + "LastUpdatedDate": "2024-01-02T03:04:05", + }, ] rc = cli_agent._cmd_list(None, 1, 10, endpoint="https://x", token="t") self.assertEqual(rc, 0) @@ -138,8 +145,8 @@ def test_download_writes_files(self): _StubClient.files = {"AGENTS.md": b"hello", "sub/x.txt": b"world"} with tempfile.TemporaryDirectory() as d: rc = cli_agent._cmd_download( - repo="user/a", local_dir=d, revision="master", - endpoint="https://x", token="t", username="user") + repo="user/a", local_dir=d, revision="master", endpoint="https://x", token="t", username="user" + ) self.assertEqual(rc, 0) self.assertEqual((Path(d) / "AGENTS.md").read_bytes(), b"hello") self.assertEqual((Path(d) / "sub" / "x.txt").read_bytes(), b"world") @@ -149,14 +156,14 @@ def test_download_missing_repo(self): _StubClient.exists = False with tempfile.TemporaryDirectory() as d: rc = cli_agent._cmd_download( - repo="user/a", local_dir=d, revision="master", - endpoint="https://x", token="t", username="user") + repo="user/a", local_dir=d, revision="master", endpoint="https://x", token="t", username="user" + ) self.assertEqual(rc, 1) def test_download_needs_owner_without_login(self): rc = cli_agent._cmd_download( - repo="a", local_dir=None, revision="master", - endpoint="https://x", token="", username="") + repo="a", local_dir=None, revision="master", endpoint="https://x", token="", username="" + ) self.assertEqual(rc, 1) @@ -171,8 +178,14 @@ def test_upload_normal_files(self): (Path(d) / "sub").mkdir() (Path(d) / "sub" / "x.txt").write_bytes(b"world") rc = cli_agent._cmd_upload( - repo="user/a", local_dir=d, revision="master", dry_run=False, - endpoint="https://x", token="t", username="user") + repo="user/a", + local_dir=d, + revision="master", + dry_run=False, + endpoint="https://x", + token="t", + username="user", + ) self.assertEqual(rc, 0) # one commit with two normal-file actions self.assertEqual(len(_StubClient.commits), 1) @@ -188,8 +201,14 @@ def test_upload_lfs_file(self): with tempfile.TemporaryDirectory() as d: (Path(d) / "model.bin").write_bytes(b"\x00\x01\x02") rc = cli_agent._cmd_upload( - repo="user/a", local_dir=d, revision="master", dry_run=False, - endpoint="https://x", token="t", username="user") + repo="user/a", + local_dir=d, + revision="master", + dry_run=False, + endpoint="https://x", + token="t", + username="user", + ) self.assertEqual(rc, 0) self.assertEqual(len(_StubClient.lfs_uploads), 1) self.assertEqual(_StubClient.lfs_uploads[0][0], "model.bin") @@ -200,8 +219,14 @@ def test_upload_creates_repo_when_absent(self): with tempfile.TemporaryDirectory() as d: (Path(d) / "AGENTS.md").write_bytes(b"hello") rc = cli_agent._cmd_upload( - repo="user/a", local_dir=d, revision="master", dry_run=False, - endpoint="https://x", token="t", username="user") + repo="user/a", + local_dir=d, + revision="master", + dry_run=False, + endpoint="https://x", + token="t", + username="user", + ) self.assertEqual(rc, 0) self.assertEqual(_StubClient.created, [("user", "a")]) @@ -210,8 +235,14 @@ def test_upload_dry_run_no_network(self): with tempfile.TemporaryDirectory() as d: (Path(d) / "AGENTS.md").write_bytes(b"hello") rc = cli_agent._cmd_upload( - repo="user/a", local_dir=d, revision="master", dry_run=True, - endpoint="https://x", token="t", username="user") + repo="user/a", + local_dir=d, + revision="master", + dry_run=True, + endpoint="https://x", + token="t", + username="user", + ) self.assertEqual(rc, 0) self.assertEqual(_StubClient.commits, []) self.assertEqual(_StubClient.lfs_uploads, []) @@ -222,9 +253,15 @@ def test_upload_passes_visibility_to_create_repo(self): with tempfile.TemporaryDirectory() as d: (Path(d) / "AGENTS.md").write_bytes(b"hello") rc = cli_agent._cmd_upload( - repo="user/a", local_dir=d, revision="master", dry_run=False, - endpoint="https://x", token="t", username="user", - visibility="private") + repo="user/a", + local_dir=d, + revision="master", + dry_run=False, + endpoint="https://x", + token="t", + username="user", + visibility="private", + ) self.assertEqual(rc, 0) self.assertEqual(_StubClient.created_visibility, ["private"]) diff --git a/tests/agent/test_anonymous_download.py b/tests/agent/test_anonymous_download.py index 11fd87a..de31508 100644 --- a/tests/agent/test_anonymous_download.py +++ b/tests/agent/test_anonymous_download.py @@ -13,6 +13,7 @@ credential loaders are patched with sentinels that must never leak into outgoing requests. """ + from __future__ import annotations import json @@ -47,8 +48,10 @@ def _mock_response(json_data=None, content=b""): def anon_api(monkeypatch): """AgentApi with an explicitly empty token and persisted creds present.""" monkeypatch.delenv("MODELSCOPE_API_TOKEN", raising=False) - with patch.object(HubConfig, "load_token", return_value=STORED), \ - patch.object(HubConfig, "load_cookies", return_value=None): + with ( + patch.object(HubConfig, "load_token", return_value=STORED), + patch.object(HubConfig, "load_cookies", return_value=None), + ): yield AgentApi(endpoint=ENDPOINT, token="", timeout=5) @@ -70,12 +73,17 @@ def test_repo_info_sends_request_without_credentials(self, anon_api): assert cookies == {} def test_list_repo_files_sends_request_without_credentials(self, anon_api): - resp = _mock_response({ - "Code": 200, "Success": True, - "Data": {"Trees": [ - {"Path": "AGENTS.md", "Type": "blob", "Sha256": "abc", "IsLfs": False}, - ]}, - }) + resp = _mock_response( + { + "Code": 200, + "Success": True, + "Data": { + "Trees": [ + {"Path": "AGENTS.md", "Type": "blob", "Sha256": "abc", "IsLfs": False}, + ] + }, + } + ) with patch.object(anon_api._openapi._session, "request", return_value=resp) as m: files = anon_api.list_repo_files("someone", "public-repo") assert files == ["AGENTS.md"] @@ -86,8 +94,7 @@ def test_list_repo_files_sends_request_without_credentials(self, anon_api): def test_download_repo_file_sends_request_without_credentials(self, anon_api): resp = _mock_response(content=b"# hello") with patch.object(anon_api._openapi._session, "request", return_value=resp) as m: - data = anon_api.download_repo_file( - "someone", "public-repo", "AGENTS.md", binary=True) + data = anon_api.download_repo_file("someone", "public-repo", "AGENTS.md", binary=True) assert data == b"# hello" auth, cookies = _sent_credentials(m) assert auth is None @@ -97,13 +104,13 @@ def test_repo_info_falls_back_to_public_probe_on_401(self, anon_api): """Server rejects anonymous /openapi metadata -> probe /api/v1 tree.""" rejected = MagicMock(status_code=401, headers={}) rejected.json.return_value = { - "success": False, "code": "InvalidAuthentication", + "success": False, + "code": "InvalidAuthentication", "message": "Invalid authentication: user not authenticated", } rejected.content = b'{"success": false}' tree_ok = _mock_response({"Code": 200, "Success": True, "Data": {"Trees": []}}) - with patch.object(anon_api._openapi._session, "request", - side_effect=[rejected, tree_ok]) as m: + with patch.object(anon_api._openapi._session, "request", side_effect=[rejected, tree_ok]) as m: info = anon_api.repo_info("someone", "public-repo") assert info == {} assert m.call_count == 2 @@ -112,23 +119,26 @@ def test_repo_info_falls_back_to_public_probe_on_401(self, anon_api): def test_repo_info_fallback_returns_none_for_missing_repo(self, anon_api): rejected = MagicMock(status_code=401, headers={}) - rejected.json.return_value = {"success": False, "code": "InvalidAuthentication", - "message": "user not authenticated"} + rejected.json.return_value = { + "success": False, + "code": "InvalidAuthentication", + "message": "user not authenticated", + } rejected.content = b'{"success": false}' missing = MagicMock(status_code=404, headers={}) - missing.json.return_value = {"Code": 10025801007, "Message": "Agent不存在", - "Success": False} + missing.json.return_value = {"Code": 10025801007, "Message": "Agent不存在", "Success": False} missing.content = b'{"Success": false}' - with patch.object(anon_api._openapi._session, "request", - side_effect=[rejected, missing]): + with patch.object(anon_api._openapi._session, "request", side_effect=[rejected, missing]): assert anon_api.repo_info("someone", "no-such-repo") is None def test_read_ops_attach_token_when_available(self, monkeypatch): """require_token=False must NOT strip credentials: with a token configured (private-repo scenario) read-only calls still send it.""" monkeypatch.delenv("MODELSCOPE_API_TOKEN", raising=False) - with patch.object(HubConfig, "load_token", return_value=None), \ - patch.object(HubConfig, "load_cookies", return_value=None): + with ( + patch.object(HubConfig, "load_token", return_value=None), + patch.object(HubConfig, "load_cookies", return_value=None), + ): api = AgentApi(endpoint=ENDPOINT, token="ms-PRIVATE-TOKEN", timeout=5) resp = _mock_response(content=b"secret file") with patch.object(api._openapi._session, "request", return_value=resp) as m: diff --git a/tests/cli/conftest.py b/tests/cli/conftest.py index 7641973..49f03d5 100644 --- a/tests/cli/conftest.py +++ b/tests/cli/conftest.py @@ -1,4 +1,5 @@ """CLI test fixtures — shared across unit and remote integration tests.""" + from __future__ import annotations import io @@ -10,7 +11,7 @@ from modelscope_hub.api import HubApi from modelscope_hub.cli.main import _build_parser, run_cmd -from modelscope_hub.types import CacheInfo, CachedRepoInfo, PagedResult, RepoInfo, UserInfo +from modelscope_hub.types import CachedRepoInfo, CacheInfo, PagedResult, RepoInfo, UserInfo # --------------------------------------------------------------------------- @@ -34,34 +35,56 @@ def mock_api(): """ api = MagicMock(spec=HubApi) api.create_repo.return_value = RepoInfo( - id=1, owner="owner", name="repo", repo_type="model", + id=1, + owner="owner", + name="repo", + repo_type="model", ) api.get_repo.return_value = RepoInfo( - id=1, owner="owner", name="repo", repo_type="model", - visibility=None, license="apache-2.0", downloads=100, likes=5, + id=1, + owner="owner", + name="repo", + repo_type="model", + visibility=None, + license="apache-2.0", + downloads=100, + likes=5, ) api.list_repos.return_value = PagedResult( items=[ - RepoInfo(id=1, owner="owner", name="model1", repo_type="model", - visibility=None, downloads=100, likes=5), + RepoInfo(id=1, owner="owner", name="model1", repo_type="model", visibility=None, downloads=100, likes=5), ], - total_count=1, page_number=1, page_size=10, + total_count=1, + page_number=1, + page_size=10, ) api.download_file.return_value = "/cache/owner/repo/file.txt" api.download_repo.return_value = "/cache/owner/repo" api.upload_folder.return_value = "commit_sha" api.whoami.return_value = UserInfo( - username="testuser", email="test@example.com", id=42, description="", + username="testuser", + email="test@example.com", + id=42, + description="", ) api.login.return_value = UserInfo( - username="testuser", email="test@example.com", id=42, + username="testuser", + email="test@example.com", + id=42, ) api.scan_cache.return_value = CacheInfo( - cache_dir="/tmp/cache", total_size=1024, - repos=[CachedRepoInfo( - repo_id="owner/repo", repo_type="model", revision="master", - nb_files=3, size_on_disk=1024, local_path="/tmp/cache/owner/repo", - )], + cache_dir="/tmp/cache", + total_size=1024, + repos=[ + CachedRepoInfo( + repo_id="owner/repo", + repo_type="model", + revision="master", + nb_files=3, + size_on_disk=1024, + local_path="/tmp/cache/owner/repo", + ) + ], ) api.clear_cache.return_value = 2048 api.list_secrets.return_value = [ @@ -69,7 +92,9 @@ def mock_api(): ] api.list_mcp_servers.return_value = PagedResult( items=[{"id": "mcp-1", "name": "weather", "status": "running", "description": "Weather MCP"}], - total_count=1, page_number=1, page_size=20, + total_count=1, + page_number=1, + page_size=20, ) api.get_mcp_server.return_value = {"id": "mcp-1", "name": "weather"} api.get_repo_logs.return_value = {"logs": ["line1", "line2"]} diff --git a/tests/cli/run_all.py b/tests/cli/run_all.py index 8239d4a..520fe91 100755 --- a/tests/cli/run_all.py +++ b/tests/cli/run_all.py @@ -9,6 +9,7 @@ This script discovers and runs every ``test_*.py`` under ``tests/cli/``, excluding tests marked ``@pytest.mark.remote`` which require live API access. """ + from __future__ import annotations import subprocess @@ -21,7 +22,9 @@ def main() -> int: cmd = [ - sys.executable, "-m", "pytest", + sys.executable, + "-m", + "pytest", str(_TESTS_DIR), "-v", "--tb=short", diff --git a/tests/cli/test_base.py b/tests/cli/test_base.py index 2bad440..5bd2d3b 100644 --- a/tests/cli/test_base.py +++ b/tests/cli/test_base.py @@ -1,4 +1,5 @@ """Tests for base.py helper functions — render_table, parse_kv_pairs, make_api, etc.""" + from __future__ import annotations from argparse import ArgumentParser, Namespace @@ -110,9 +111,7 @@ def test_make_api_with_endpoint(self): args = Namespace(token=None, endpoint="https://custom.endpoint.com") with patch("modelscope_hub.cli.base.HubApi") as mock_hub: make_api(args) - mock_hub.assert_called_once_with( - token=None, endpoint="https://custom.endpoint.com" - ) + mock_hub.assert_called_once_with(token=None, endpoint="https://custom.endpoint.com") class TestAddRepoTypeArg: diff --git a/tests/cli/test_cache.py b/tests/cli/test_cache.py index 112bb17..02e2a61 100644 --- a/tests/cli/test_cache.py +++ b/tests/cli/test_cache.py @@ -313,7 +313,6 @@ def test_mismatch_exits_nonzero(self, parser, mock_api, capsys): assert exc_info.value.code == 1 assert "weights.bin" in capsys.readouterr().err - def test_missing_and_extra_warnings_include_bounded_paths(self, parser, mock_api, capsys): mock_api.verify_cache.return_value = CacheVerification( revision="master", diff --git a/tests/cli/test_compat.py b/tests/cli/test_compat.py index 36c5900..c743a29 100644 --- a/tests/cli/test_compat.py +++ b/tests/cli/test_compat.py @@ -20,10 +20,15 @@ class TestDownloadLegacyEdgeCases: def test_legacy_dataset_with_local_dir(self, parser): """ms download --dataset owner/repo --local_dir ./temp (regression test)""" - args = parser.parse_args([ - "download", "--dataset", "wangxingjun778/self_cog_data", - "--local_dir", "./temp", - ]) + args = parser.parse_args( + [ + "download", + "--dataset", + "wangxingjun778/self_cog_data", + "--local_dir", + "./temp", + ] + ) assert args.dataset == "wangxingjun778/self_cog_data" assert args.local_dir_legacy == "./temp" @@ -49,9 +54,15 @@ def test_dataset_to_repo_id(self, parser): assert args.repo_type == "dataset" def test_local_dir_legacy_merged(self, parser): - args = parser.parse_args([ - "download", "--model", "owner/repo", "--local_dir", "/tmp/out", - ]) + args = parser.parse_args( + [ + "download", + "--model", + "owner/repo", + "--local_dir", + "/tmp/out", + ] + ) with warnings.catch_warnings(): warnings.simplefilter("ignore", DeprecationWarning) normalize_download_args(args) @@ -120,11 +131,14 @@ def test_version_short(self, parser): # Cross-cutting backward compat: --repo_type (underscore) in multiple commands # --------------------------------------------------------------------------- class TestRepoTypeUnderscore: - @pytest.mark.parametrize("cmd,expected_type", [ - (["info", "o/r", "--repo_type", "dataset"], "dataset"), - (["list", "--repo_type", "model"], "model"), - (["delete", "o/r", "--repo_type", "model"], "model"), - ]) + @pytest.mark.parametrize( + "cmd,expected_type", + [ + (["info", "o/r", "--repo_type", "dataset"], "dataset"), + (["list", "--repo_type", "model"], "model"), + (["delete", "o/r", "--repo_type", "model"], "model"), + ], + ) def test_repo_type_underscore_in_all_commands(self, parser, cmd, expected_type): """--repo_type (underscore) works in info, list, delete.""" args = parser.parse_args(cmd) diff --git a/tests/cli/test_compat_cache_dir.py b/tests/cli/test_compat_cache_dir.py index a036ae2..ed45f47 100644 --- a/tests/cli/test_compat_cache_dir.py +++ b/tests/cli/test_compat_cache_dir.py @@ -84,7 +84,7 @@ def test_cache_dir_passed_through(self, MockHubApi): mock_api.download_repo.return_value = "/tmp/cache/models/damo--bert/snapshots/master" MockHubApi.return_value = mock_api - result = snapshot_download(model_id="damo/bert", cache_dir="/tmp/cache") + snapshot_download(model_id="damo/bert", cache_dir="/tmp/cache") mock_api.download_repo.assert_called_once() call_kwargs = mock_api.download_repo.call_args[1] @@ -101,7 +101,7 @@ def test_local_dir_explicit_not_overridden(self, MockHubApi): mock_api.download_repo.return_value = "/custom/dir" MockHubApi.return_value = mock_api - result = snapshot_download(model_id="damo/bert", local_dir="/custom/dir") + snapshot_download(model_id="damo/bert", local_dir="/custom/dir") mock_api.download_repo.assert_called_once() call_kwargs = mock_api.download_repo.call_args[1] @@ -117,8 +117,9 @@ def test_dataset_snapshot_download_cache_dir_passthrough(self, MockHubApi): mock_api.download_repo.return_value = "/data/hub/datasets/my_org--dataset1/snapshots/master" MockHubApi.return_value = mock_api - result = dataset_snapshot_download( - dataset_id="my_org/dataset1", cache_dir="/data/hub", + dataset_snapshot_download( + dataset_id="my_org/dataset1", + cache_dir="/data/hub", ) mock_api.download_repo.assert_called_once() @@ -144,7 +145,7 @@ def test_cache_dir_passed_through(self, MockHubApi): mock_api.download_file.return_value = "/data/hub/models/qwen--chat/snapshots/master/model.bin" MockHubApi.return_value = mock_api - result = model_file_download("qwen/chat", "model.bin", cache_dir="/data/hub") + model_file_download("qwen/chat", "model.bin", cache_dir="/data/hub") mock_api.download_file.assert_called_once() call_kwargs = mock_api.download_file.call_args[1] @@ -161,8 +162,10 @@ def test_local_dir_explicit_passthrough(self, MockHubApi): mock_api.download_file.return_value = "/my/dir/model.bin" MockHubApi.return_value = mock_api - result = model_file_download( - "qwen/chat", "model.bin", local_dir="/my/dir", + model_file_download( + "qwen/chat", + "model.bin", + local_dir="/my/dir", ) mock_api.download_file.assert_called_once() @@ -179,8 +182,10 @@ def test_dataset_file_download_cache_dir_passthrough(self, MockHubApi): mock_api.download_file.return_value = "/data/hub/datasets/org--ds/snapshots/master/train.csv" MockHubApi.return_value = mock_api - result = dataset_file_download( - "org/ds", "train.csv", cache_dir="/data/hub", + dataset_file_download( + "org/ds", + "train.csv", + cache_dir="/data/hub", ) mock_api.download_file.assert_called_once() @@ -206,7 +211,7 @@ def test_no_args_uses_standard_cache(self, MockHubApi, tmp_path): mock_api.download_file.return_value = "/default/cache/models/owner--name/snapshots/master/README.md" MockHubApi.return_value = mock_api - result = model_file_download("owner/name", "README.md") + model_file_download("owner/name", "README.md") mock_api.download_file.assert_called_once() call_kwargs = mock_api.download_file.call_args[1] @@ -223,7 +228,7 @@ def test_snapshot_no_args_uses_standard_cache(self, MockHubApi): mock_api.download_repo.return_value = "/default/cache/models/org--model/snapshots/master" MockHubApi.return_value = mock_api - result = snapshot_download(model_id="org/model") + snapshot_download(model_id="org/model") mock_api.download_repo.assert_called_once() call_kwargs = mock_api.download_repo.call_args[1] diff --git a/tests/cli/test_compat_revision.py b/tests/cli/test_compat_revision.py index a3a05c7..ed55f9c 100644 --- a/tests/cli/test_compat_revision.py +++ b/tests/cli/test_compat_revision.py @@ -123,14 +123,17 @@ class TestDevMode: def test_dev_mode_defaults_to_master(self): api = _make_api([_make_rev("master")], [_make_rev("v1.0")]) detail = api.get_valid_revision_detail( - "o/m", release_timestamp=self.FAR_FUTURE, + "o/m", + release_timestamp=self.FAR_FUTURE, ) assert detail["Revision"] == "master" def test_dev_mode_explicit_tag(self): api = _make_api([_make_rev("master")], [_make_rev("v2.0")]) detail = api.get_valid_revision_detail( - "o/m", revision="v2.0", release_timestamp=self.FAR_FUTURE, + "o/m", + revision="v2.0", + release_timestamp=self.FAR_FUTURE, ) assert detail["Revision"] == "v2.0" @@ -138,7 +141,9 @@ def test_dev_mode_nonexistent_raises(self): api = _make_api([_make_rev("master")], []) with pytest.raises(NotExistError): api.get_valid_revision_detail( - "o/m", revision="nope", release_timestamp=self.FAR_FUTURE, + "o/m", + revision="nope", + release_timestamp=self.FAR_FUTURE, ) @@ -154,21 +159,26 @@ def test_explicit_branch_returns_immediately(self): [_make_rev("v1.0", 500)], ) detail = api.get_valid_revision_detail( - "o/m", revision="dev", release_timestamp=self.RELEASE_TS, + "o/m", + revision="dev", + release_timestamp=self.RELEASE_TS, ) assert detail["Revision"] == "dev" def test_no_tags_defaults_to_master(self): api = _make_api([_make_rev("master", 100)], []) detail = api.get_valid_revision_detail( - "o/m", release_timestamp=self.RELEASE_TS, + "o/m", + release_timestamp=self.RELEASE_TS, ) assert detail["Revision"] == "master" def test_no_tags_explicit_master(self): api = _make_api([_make_rev("master", 100)], []) detail = api.get_valid_revision_detail( - "o/m", revision="master", release_timestamp=self.RELEASE_TS, + "o/m", + revision="master", + release_timestamp=self.RELEASE_TS, ) assert detail["Revision"] == "master" @@ -176,7 +186,9 @@ def test_no_tags_explicit_nonexistent_raises(self): api = _make_api([_make_rev("master")], []) with pytest.raises(NotExistError): api.get_valid_revision_detail( - "o/m", revision="v1.0", release_timestamp=self.RELEASE_TS, + "o/m", + revision="v1.0", + release_timestamp=self.RELEASE_TS, ) def test_auto_selects_latest_tag_before_release(self): @@ -189,7 +201,8 @@ def test_auto_selects_latest_tag_before_release(self): ], ) detail = api.get_valid_revision_detail( - "o/m", release_timestamp=self.RELEASE_TS, + "o/m", + release_timestamp=self.RELEASE_TS, ) # v3.0 (1500) is the newest with CreatedAt <= 2000 assert detail["Revision"] == "v3.0" @@ -200,7 +213,8 @@ def test_no_tag_before_release_falls_back_to_master(self): [_make_rev("v1.0", 3000)], ) detail = api.get_valid_revision_detail( - "o/m", release_timestamp=self.RELEASE_TS, + "o/m", + release_timestamp=self.RELEASE_TS, ) assert detail["Revision"] == "master" @@ -210,7 +224,9 @@ def test_explicit_valid_tag(self): [_make_rev("v1.0", 500), _make_rev("v2.0", 1000)], ) detail = api.get_valid_revision_detail( - "o/m", revision="v1.0", release_timestamp=self.RELEASE_TS, + "o/m", + revision="v1.0", + release_timestamp=self.RELEASE_TS, ) assert detail["Revision"] == "v1.0" @@ -221,7 +237,9 @@ def test_explicit_invalid_tag_raises(self): ) with pytest.raises(NotExistError, match="valid tags"): api.get_valid_revision_detail( - "o/m", revision="v999", release_timestamp=self.RELEASE_TS, + "o/m", + revision="v999", + release_timestamp=self.RELEASE_TS, ) def test_explicit_master_with_tags_allowed(self): @@ -231,6 +249,8 @@ def test_explicit_master_with_tags_allowed(self): [_make_rev("v1.0", 500)], ) detail = api.get_valid_revision_detail( - "o/m", revision="master", release_timestamp=self.RELEASE_TS, + "o/m", + revision="master", + release_timestamp=self.RELEASE_TS, ) assert detail["Revision"] == "master" diff --git a/tests/cli/test_deploy.py b/tests/cli/test_deploy.py index bed20bf..fe855da 100644 --- a/tests/cli/test_deploy.py +++ b/tests/cli/test_deploy.py @@ -5,6 +5,7 @@ - Execution tests: mock HubApi to verify command logic - Remote tests: real API lifecycle (existing) """ + from __future__ import annotations import warnings @@ -37,9 +38,16 @@ def test_repo_type_default_studio(self, parser): assert args.repo_type == "studio" def test_subcmd_token_endpoint(self, parser): - args = parser.parse_args([ - "deploy", "o/r", "--token", "tk", "--endpoint", "https://x.cn", - ]) + args = parser.parse_args( + [ + "deploy", + "o/r", + "--token", + "tk", + "--endpoint", + "https://x.cn", + ] + ) assert args.subcmd_token == "tk" assert args.subcmd_endpoint == "https://x.cn" @@ -124,13 +132,20 @@ def test_keyword_default_none(self, parser): assert args.keyword is None def test_all_options_combined(self, parser): - args = parser.parse_args([ - "logs", "org/demo", - "--log-type", "build", - "--page", "2", - "--page-size", "50", - "--keyword", "Exception", - ]) + args = parser.parse_args( + [ + "logs", + "org/demo", + "--log-type", + "build", + "--page", + "2", + "--page-size", + "50", + "--keyword", + "Exception", + ] + ) assert args.repo_id == "org/demo" assert args.log_type == "build" assert args.page_num == 2 @@ -236,8 +251,12 @@ def test_logs_with_list_payload(self, parser, mock_api, capsys): with patch("modelscope_hub.cli.deploy.make_api", return_value=mock_api): LogsCommand(args).execute() mock_api.get_repo_logs.assert_called_once_with( - "org/demo", "studio", - log_type="run", page_num=1, page_size=100, keyword=None, + "org/demo", + "studio", + log_type="run", + page_num=1, + page_size=100, + keyword=None, ) out = capsys.readouterr().out assert "line1" in out @@ -279,7 +298,9 @@ def test_single_setting(self, parser, mock_api, capsys): with patch("modelscope_hub.cli.deploy.make_api", return_value=mock_api): SettingsCommand(args).execute() mock_api.update_repo_settings.assert_called_once_with( - "org/demo", "studio", cpu="4", + "org/demo", + "studio", + cpu="4", ) out = capsys.readouterr().out assert "Updated 1 setting" in out @@ -289,15 +310,24 @@ def test_multiple_settings(self, parser, mock_api, capsys): with patch("modelscope_hub.cli.deploy.make_api", return_value=mock_api): SettingsCommand(args).execute() mock_api.update_repo_settings.assert_called_once_with( - "org/demo", "studio", cpu="4", memory="8192", + "org/demo", + "studio", + cpu="4", + memory="8192", ) out = capsys.readouterr().out assert "Updated 2 setting" in out def test_settings_skill_type(self, parser, mock_api, capsys): - args = parser.parse_args([ - "settings", "org/skill1", "timeout=30", "--repo-type", "skill", - ]) + args = parser.parse_args( + [ + "settings", + "org/skill1", + "timeout=30", + "--repo-type", + "skill", + ] + ) with patch("modelscope_hub.cli.deploy.make_api", return_value=mock_api): SettingsCommand(args).execute() assert mock_api.update_repo_settings.call_args.args[1] == "skill" diff --git a/tests/cli/test_download.py b/tests/cli/test_download.py index 4b5cf3d..f442436 100644 --- a/tests/cli/test_download.py +++ b/tests/cli/test_download.py @@ -5,17 +5,17 @@ - Execution tests: mock HubApi for file/snapshot download logic - Remote tests: real API file download (existing) """ + from __future__ import annotations import warnings +from pathlib import Path from unittest.mock import patch import pytest from modelscope_hub.cli.download import DownloadCommand -from pathlib import Path - from .conftest import run_cli @@ -85,9 +85,16 @@ def test_include_multi(self, parser): assert args.allow_patterns == ["*.bin", "*.json"] def test_include_repeated(self, parser): - args = parser.parse_args([ - "download", "o/r", "--include", "*.bin", "--include", "*.json", - ]) + args = parser.parse_args( + [ + "download", + "o/r", + "--include", + "*.bin", + "--include", + "*.json", + ] + ) assert "*.bin" in args.allow_patterns assert "*.json" in args.allow_patterns @@ -100,26 +107,42 @@ def test_exclude_multi(self, parser): assert args.ignore_patterns == ["*.bin", "*.gguf"] def test_include_and_exclude(self, parser): - args = parser.parse_args([ - "download", "o/r", - "--include", "*.safetensors", - "--exclude", "*.bin", "*.gguf", - ]) + args = parser.parse_args( + [ + "download", + "o/r", + "--include", + "*.safetensors", + "--exclude", + "*.bin", + "*.gguf", + ] + ) assert args.allow_patterns == ["*.safetensors"] assert args.ignore_patterns == ["*.bin", "*.gguf"] def test_all_options_combined(self, parser): - args = parser.parse_args([ - "download", "Qwen/Qwen3-0.6B", - "--repo-type", "model", - "--revision", "main", - "--cache-dir", "/cache", - "--local-dir", "./out", - "--max-workers", "16", - "--include", "*.safetensors", - "--exclude", "*.bin", - "--force", - ]) + args = parser.parse_args( + [ + "download", + "Qwen/Qwen3-0.6B", + "--repo-type", + "model", + "--revision", + "main", + "--cache-dir", + "/cache", + "--local-dir", + "./out", + "--max-workers", + "16", + "--include", + "*.safetensors", + "--exclude", + "*.bin", + "--force", + ] + ) assert args.repo_id == "Qwen/Qwen3-0.6B" assert args.repo_type == "model" assert args.revision == "main" @@ -161,10 +184,17 @@ def test_legacy_cache_dir_underscore(self, parser): assert args.cache_dir_legacy == "/cache" def test_subcmd_token_endpoint(self, parser): - args = parser.parse_args([ - "download", "--model", "o/r", - "--token", "ms-xxx", "--endpoint", "https://custom.cn", - ]) + args = parser.parse_args( + [ + "download", + "--model", + "o/r", + "--token", + "ms-xxx", + "--endpoint", + "https://custom.cn", + ] + ) assert args.subcmd_token == "ms-xxx" assert args.subcmd_endpoint == "https://custom.cn" @@ -185,9 +215,15 @@ def _patch_download_api(self, mock_api): ) def test_single_file(self, parser, mock_api, capsys): - args = parser.parse_args([ - "download", "owner/repo", "config.json", "--cache-dir", "/tmp/cache", - ]) + args = parser.parse_args( + [ + "download", + "owner/repo", + "config.json", + "--cache-dir", + "/tmp/cache", + ] + ) p1, p2 = self._patch_download_api(mock_api) with p1, p2: DownloadCommand(args).execute() @@ -199,9 +235,14 @@ def test_single_file(self, parser, mock_api, capsys): assert "config.json" in out def test_multiple_files(self, parser, mock_api, capsys): - args = parser.parse_args([ - "download", "owner/repo", "a.bin", "b.json", - ]) + args = parser.parse_args( + [ + "download", + "owner/repo", + "a.bin", + "b.json", + ] + ) p1, p2 = self._patch_download_api(mock_api) with p1, p2: DownloadCommand(args).execute() @@ -217,11 +258,16 @@ def test_snapshot(self, parser, mock_api, capsys): assert "Snapshot ready" in out def test_snapshot_with_patterns(self, parser, mock_api, capsys): - args = parser.parse_args([ - "download", "owner/repo", - "--include", "*.safetensors", - "--exclude", "*.bin", - ]) + args = parser.parse_args( + [ + "download", + "owner/repo", + "--include", + "*.safetensors", + "--exclude", + "*.bin", + ] + ) p1, p2 = self._patch_download_api(mock_api) with p1, p2: DownloadCommand(args).execute() @@ -237,9 +283,14 @@ def test_force_forwarded(self, parser, mock_api, capsys): assert mock_api.download_file.call_args.kwargs["force"] is True def test_dataset_repo_type(self, parser, mock_api, capsys): - args = parser.parse_args([ - "download", "org/data", "--repo-type", "dataset", - ]) + args = parser.parse_args( + [ + "download", + "org/data", + "--repo-type", + "dataset", + ] + ) p1, p2 = self._patch_download_api(mock_api) with p1, p2: DownloadCommand(args).execute() @@ -325,15 +376,17 @@ def setup_repo(self, api, test_owner, test_endpoint, repo_name): api.legacy.create_commit( repo_id=cls.repo_id, repo_type="model", - operations=[{ - "action": "create", - "path": "test_data.txt", - "type": "normal", - "size": len(file_bytes), - "sha256": "", - "content": content_b64, - "encoding": "base64", - }], + operations=[ + { + "action": "create", + "path": "test_data.txt", + "type": "normal", + "size": len(file_bytes), + "sha256": "", + "content": content_b64, + "encoding": "base64", + } + ], commit_message="Add test file", revision="master", ) diff --git a/tests/cli/test_login.py b/tests/cli/test_login.py index 09095fa..0dc939f 100644 --- a/tests/cli/test_login.py +++ b/tests/cli/test_login.py @@ -5,6 +5,7 @@ - Execution tests: mock HubApi for login/whoami logic - Remote tests: real API (existing) """ + from __future__ import annotations from unittest.mock import patch @@ -44,9 +45,15 @@ def test_no_args(self, parser): assert hasattr(args, "_command") def test_subcmd_token_endpoint(self, parser): - args = parser.parse_args([ - "whoami", "--token", "my-tok", "--endpoint", "https://x.cn", - ]) + args = parser.parse_args( + [ + "whoami", + "--token", + "my-tok", + "--endpoint", + "https://x.cn", + ] + ) assert args.subcmd_token == "my-tok" assert args.subcmd_endpoint == "https://x.cn" @@ -123,9 +130,15 @@ def test_login_eof_exits_130(self, parser, mock_api): assert exc_info.value.code == 130 def test_login_subcmd_endpoint_merged(self, parser, mock_api, capsys): - args = parser.parse_args([ - "login", "--token", "tok", "--endpoint", "https://custom.cn", - ]) + args = parser.parse_args( + [ + "login", + "--token", + "tok", + "--endpoint", + "https://custom.cn", + ] + ) with patch("modelscope_hub.cli.login.make_api", return_value=mock_api): LoginCommand(args).execute() assert args.endpoint == "https://custom.cn" diff --git a/tests/cli/test_main.py b/tests/cli/test_main.py index a43b801..4abce64 100644 --- a/tests/cli/test_main.py +++ b/tests/cli/test_main.py @@ -1,4 +1,5 @@ """Tests for CLI entry point, global parameters, exception handling, and version.""" + from __future__ import annotations import logging @@ -6,7 +7,6 @@ import pytest -from modelscope_hub import __version__ from modelscope_hub.cli.main import run_cmd from modelscope_hub.errors import HubError, InvalidParameter, NetworkError, NotSupportedError @@ -70,9 +70,16 @@ def test_verbose_default_false(self, parser): assert args.verbose is False def test_global_flags_before_subcommand(self, parser): - args = parser.parse_args([ - "--token", "tok", "--endpoint", "https://x.cn", "-v", "whoami", - ]) + args = parser.parse_args( + [ + "--token", + "tok", + "--endpoint", + "https://x.cn", + "-v", + "whoami", + ] + ) assert args.token == "tok" assert args.endpoint == "https://x.cn" assert args.verbose is True @@ -104,9 +111,7 @@ def test_invalid_parameter_exits_2(self): def test_not_supported_error_exits_2(self): with patch("modelscope_hub.cli.login.make_api") as mock_make: - mock_make.return_value.whoami.side_effect = NotSupportedError( - "not supported", suggestion="use Y" - ) + mock_make.return_value.whoami.side_effect = NotSupportedError("not supported", suggestion="use Y") code, out, err = run_cli(["whoami"], token="fake") assert code == 2 assert "not supported" in err diff --git a/tests/cli/test_mcp.py b/tests/cli/test_mcp.py index a5a114f..06be932 100644 --- a/tests/cli/test_mcp.py +++ b/tests/cli/test_mcp.py @@ -5,6 +5,7 @@ - Execution tests: mock HubApi for MCP server operations - Remote tests: real API (existing) """ + from __future__ import annotations from unittest.mock import patch @@ -28,9 +29,16 @@ def test_basic(self, parser): assert hasattr(args, "_mcp_leaf") def test_subcmd_token_endpoint(self, parser): - args = parser.parse_args([ - "mcp", "list", "--token", "tk", "--endpoint", "https://x.cn", - ]) + args = parser.parse_args( + [ + "mcp", + "list", + "--token", + "tk", + "--endpoint", + "https://x.cn", + ] + ) assert args.subcmd_token == "tk" assert args.subcmd_endpoint == "https://x.cn" @@ -56,12 +64,18 @@ def test_page_defaults(self, parser): assert args.page_size == 20 def test_all_options(self, parser): - args = parser.parse_args([ - "mcp", "list", - "--search", "test", - "--page", "2", - "--page-size", "50", - ]) + args = parser.parse_args( + [ + "mcp", + "list", + "--search", + "test", + "--page", + "2", + "--page-size", + "50", + ] + ) assert args.search == "test" assert args.page_number == 2 assert args.page_size == 50 @@ -79,9 +93,17 @@ def test_missing_server_id_exits(self, parser): parser.parse_args(["mcp", "info"]) def test_subcmd_token_endpoint(self, parser): - args = parser.parse_args([ - "mcp", "info", "org/srv", "--token", "tk", "--endpoint", "https://x.cn", - ]) + args = parser.parse_args( + [ + "mcp", + "info", + "org/srv", + "--token", + "tk", + "--endpoint", + "https://x.cn", + ] + ) assert args.subcmd_token == "tk" assert args.subcmd_endpoint == "https://x.cn" @@ -98,9 +120,17 @@ def test_missing_server_id_exits(self, parser): parser.parse_args(["mcp", "deploy"]) def test_subcmd_token_endpoint(self, parser): - args = parser.parse_args([ - "mcp", "deploy", "org/srv", "--token", "tk", "--endpoint", "https://x.cn", - ]) + args = parser.parse_args( + [ + "mcp", + "deploy", + "org/srv", + "--token", + "tk", + "--endpoint", + "https://x.cn", + ] + ) assert args.subcmd_token == "tk" @@ -116,9 +146,17 @@ def test_missing_server_id_exits(self, parser): parser.parse_args(["mcp", "undeploy"]) def test_subcmd_token_endpoint(self, parser): - args = parser.parse_args([ - "mcp", "undeploy", "org/srv", "--token", "tk", "--endpoint", "https://x.cn", - ]) + args = parser.parse_args( + [ + "mcp", + "undeploy", + "org/srv", + "--token", + "tk", + "--endpoint", + "https://x.cn", + ] + ) assert args.subcmd_token == "tk" @@ -132,14 +170,19 @@ def test_list_with_results(self, parser, mock_api, capsys): with patch("modelscope_hub.cli.mcp.make_api", return_value=mock_api): _McpList(args).execute() mock_api.list_mcp_servers.assert_called_once_with( - search=None, page_number=1, page_size=20, + search=None, + page_number=1, + page_size=20, ) out = capsys.readouterr().out assert "weather" in out def test_list_empty(self, parser, mock_api, capsys): mock_api.list_mcp_servers.return_value = PagedResult( - items=[], total_count=0, page_number=1, page_size=20, + items=[], + total_count=0, + page_number=1, + page_size=20, ) args = parser.parse_args(["mcp", "list"]) with patch("modelscope_hub.cli.mcp.make_api", return_value=mock_api): @@ -179,10 +222,61 @@ def test_deploy(self, parser, mock_api, capsys): args = parser.parse_args(["mcp", "deploy", "org/weather-mcp"]) with patch("modelscope_hub.cli.mcp.make_api", return_value=mock_api): _McpDeploy(args).execute() - mock_api.deploy_mcp_server.assert_called_once_with("org/weather-mcp") + # No CLI options -> the payload is left to the API layer defaults. + mock_api.deploy_mcp_server.assert_called_once_with("org/weather-mcp", payload=None) out = capsys.readouterr().out assert "Deploy requested" in out + def test_deploy_with_options(self, parser, mock_api, capsys): + args = parser.parse_args( + [ + "mcp", + "deploy", + "org/weather-mcp", + "--transport-type", + "streamable_http", + "--expiration-minutes", + "60", + "--auth-check", + "--env", + "API_KEY=abc", + "--env", + "REGION=cn", + ] + ) + with patch("modelscope_hub.cli.mcp.make_api", return_value=mock_api): + _McpDeploy(args).execute() + mock_api.deploy_mcp_server.assert_called_once_with( + "org/weather-mcp", + payload={ + "transport_type": "streamable_http", + "expiration_minutes": 60, + "auth_check": True, + "env_info": {"API_KEY": "abc", "REGION": "cn"}, + }, + ) + out = capsys.readouterr().out + assert "Deploy requested" in out + + def test_deploy_rejects_invalid_transport(self, parser): + with pytest.raises(SystemExit): + parser.parse_args(["mcp", "deploy", "org/weather-mcp", "--transport-type", "bogus"]) + + def test_deploy_prints_operational_url(self, parser, mock_api, capsys): + mock_api.deploy_mcp_server.return_value = { + "id": "platform-pool", + "url": "https://mcp.api-inference.modelscope.net/xxxx/mcp", + "transport_type": "sse", + "expiration": "2025-10-01 21:00:00", + "auth_required": False, + } + args = parser.parse_args(["mcp", "deploy", "org/weather-mcp"]) + with patch("modelscope_hub.cli.mcp.make_api", return_value=mock_api): + _McpDeploy(args).execute() + out = capsys.readouterr().out + assert "https://mcp.api-inference.modelscope.net/xxxx/mcp" in out + assert "expires: 2025-10-01 21:00:00" in out + @pytest.mark.mock_only class TestMcpUndeployExecute: @@ -209,7 +303,7 @@ def test_list_mcp_servers(self, test_token, test_endpoint): token=test_token, endpoint=test_endpoint, ) - print(f"\n** [mcp list]") + print("\n** [mcp list]") print(f"** exit_code={exit_code}, out={out[:300]!r}, err={err!r}") assert exit_code == 0 assert "mcp" in out.lower() or "no MCP servers found" in out or "id" in out.lower() @@ -221,6 +315,6 @@ def test_list_mcp_with_search(self, test_token, test_endpoint): token=test_token, endpoint=test_endpoint, ) - print(f"\n** [mcp list --search test]") + print("\n** [mcp list --search test]") print(f"** exit_code={exit_code}, out={out[:300]!r}, err={err!r}") assert exit_code == 0 diff --git a/tests/cli/test_openapi.py b/tests/cli/test_openapi.py index d3250cf..1625da0 100644 --- a/tests/cli/test_openapi.py +++ b/tests/cli/test_openapi.py @@ -2,6 +2,7 @@ Covers fixes from audit items 2,3,4,5,6,8,10 and Section III risks. """ + from __future__ import annotations from unittest.mock import MagicMock, patch @@ -9,7 +10,7 @@ import pytest import requests -from modelscope_hub._openapi import OpenAPIClient, _RETRYABLE_POST_PATHS +from modelscope_hub._openapi import _RETRYABLE_POST_PATHS, OpenAPIClient from modelscope_hub.api import HubApi from modelscope_hub.config import HubConfig from modelscope_hub.errors import InvalidParameter, RateLimitError, ServerError @@ -148,6 +149,7 @@ def test_no_json_body_sent(self, client): class TestGetStudioAuth: def test_requires_token_raises_without_token(self): from modelscope_hub.errors import AuthenticationError + config = HubConfig(token="placeholder", endpoint="https://modelscope.cn") config.token = None client = OpenAPIClient(config) @@ -244,14 +246,14 @@ def test_deploy_studio_retried_on_server_error(self, client): error_resp = _mock_response(status_code=500, json_data={"message": "Internal error"}) success_resp = _mock_response(status_code=200, json_data={"success": True, "data": {"status": "deploying"}}) with patch.object(client._session, "request", side_effect=[error_resp, success_resp]) as mock_req: - result = client.deploy_studio("org", "demo") + client.deploy_studio("org", "demo") assert mock_req.call_count == 2 def test_stop_studio_retried_on_server_error(self, client): error_resp = _mock_response(status_code=500, json_data={"message": "Internal error"}) success_resp = _mock_response(status_code=200, json_data={"success": True, "data": {"status": "stopped"}}) with patch.object(client._session, "request", side_effect=[error_resp, success_resp]) as mock_req: - result = client.stop_studio("org", "demo") + client.stop_studio("org", "demo") assert mock_req.call_count == 2 def test_create_skill_not_retried(self, client): @@ -265,10 +267,38 @@ def test_deploy_mcp_server_retried(self, client): error_resp = _mock_response(status_code=500, json_data={"message": "Internal error"}) success_resp = _mock_response(status_code=200, json_data={"success": True, "data": {"status": "running"}}) with patch.object(client._session, "request", side_effect=[error_resp, success_resp]) as mock_req: - result = client.deploy_mcp_server("123") + client.deploy_mcp_server("123") assert mock_req.call_count == 2 +# ================================================================== +# deploy_mcp_server payload hygiene: None values are dropped before +# the default transport is applied, so ``transport_type=None`` can +# never reach the wire. +# ================================================================== +class TestDeployMcpServerPayload: + def test_default_transport_applied(self, client): + resp = _mock_response(json_data={"success": True, "data": {}}) + with patch.object(client._session, "request", return_value=resp) as mock_req: + client.deploy_mcp_server("123") + assert mock_req.call_args.kwargs["json"] == {"transport_type": "sse"} + + def test_explicit_none_replaced_by_default(self, client): + resp = _mock_response(json_data={"success": True, "data": {}}) + with patch.object(client._session, "request", return_value=resp) as mock_req: + client.deploy_mcp_server("123", {"transport_type": None, "expiration_minutes": None}) + assert mock_req.call_args.kwargs["json"] == {"transport_type": "sse"} + + def test_caller_payload_preserved(self, client): + resp = _mock_response(json_data={"success": True, "data": {}}) + payload = {"transport_type": "streamable_http", "expiration_minutes": -1, "env_info": {"K": "v"}} + with patch.object(client._session, "request", return_value=resp) as mock_req: + client.deploy_mcp_server("123", payload) + assert mock_req.call_args.kwargs["json"] == payload + # The client works on a copy; the caller's dict is untouched. + assert "transport_type" in payload and payload["expiration_minutes"] == -1 + + # ================================================================== # Rate-limit / commit-lock-busy retry on a non-idempotent POST. # A 429 means the server rejected the request WITHOUT processing it, @@ -284,9 +314,10 @@ def test_commit_style_post_retried_on_rate_limit(self, client): json_data={"message": "commit lock busy, please try again"}, ) ok = _mock_response(status_code=200, json_data={"success": True, "data": {}}) - with patch("modelscope_hub._openapi.time.sleep"), \ - patch.object(client._session, "request", - side_effect=[busy, busy, ok]) as mock_req: + with ( + patch("modelscope_hub._openapi.time.sleep"), + patch.object(client._session, "request", side_effect=[busy, busy, ok]) as mock_req, + ): client.request("POST", url=self._COMMIT_URL, json_body={"actions": []}) # two 429s then success -> three calls total (i.e. it retried). assert mock_req.call_count == 3 @@ -296,8 +327,10 @@ def test_commit_style_post_gives_up_after_max_retries(self, client): status_code=429, json_data={"message": "commit lock busy, please try again"}, ) - with patch("modelscope_hub._openapi.time.sleep"), \ - patch.object(client._session, "request", return_value=busy) as mock_req: + with ( + patch("modelscope_hub._openapi.time.sleep"), + patch.object(client._session, "request", return_value=busy) as mock_req, + ): with pytest.raises(RateLimitError): client.request("POST", url=self._COMMIT_URL, json_body={"actions": []}) # Exhausts all attempts rather than failing on the first 429. @@ -306,8 +339,10 @@ def test_commit_style_post_gives_up_after_max_retries(self, client): def test_bare_post_not_retried_on_plain_400(self, client): """A non-rate-limit 400 on a bare url= POST is still NOT retried.""" bad = _mock_response(status_code=400, json_data={"message": "invalid parameter"}) - with patch("modelscope_hub._openapi.time.sleep"), \ - patch.object(client._session, "request", return_value=bad) as mock_req: + with ( + patch("modelscope_hub._openapi.time.sleep"), + patch.object(client._session, "request", return_value=bad) as mock_req, + ): with pytest.raises(InvalidParameter): client.request("POST", url=self._COMMIT_URL, json_body={"actions": []}) assert mock_req.call_count == 1 @@ -325,17 +360,20 @@ def test_same_host_absolute_url_gets_auth(self, client): client.request("POST", url=url, json_body={}) call_kwargs = mock_req.call_args.kwargs assert call_kwargs["headers"].get("Authorization") == "Bearer test-token" - assert call_kwargs["cookies"] == { - "m_session_id": "test-token", "modelscope_session": "test-token"} + assert call_kwargs["cookies"] == {"m_session_id": "test-token", "modelscope_session": "test-token"} def test_foreign_host_absolute_url_strips_auth_and_cookies(self, client): resp = _mock_response() url = "https://oss-cn-hangzhou.aliyuncs.com/bucket/obj?sig=abc" with patch.object(client._session, "request", return_value=resp) as mock_req: client.request( - "PUT", url=url, data=b"blob", + "PUT", + url=url, + data=b"blob", headers={"Content-Type": "application/octet-stream"}, - require_token=False, unwrap=False) + require_token=False, + unwrap=False, + ) call_kwargs = mock_req.call_args.kwargs assert "Authorization" not in call_kwargs["headers"] assert call_kwargs["cookies"] == {} diff --git a/tests/cli/test_repo.py b/tests/cli/test_repo.py index 15b2cb4..3e45300 100644 --- a/tests/cli/test_repo.py +++ b/tests/cli/test_repo.py @@ -5,6 +5,7 @@ - Execution tests: mock HubApi to verify command logic without network - Remote tests: real API lifecycle (existing, kept as-is) """ + from __future__ import annotations import warnings @@ -45,51 +46,35 @@ def test_invalid_repo_type_rejected(self, parser): @pytest.mark.parametrize("vis", ["public", "private", "internal"]) def test_visibility_choices(self, parser, vis): - args = parser.parse_args( - ["create", "o/r", "--repo-type", "model", "--visibility", vis] - ) + args = parser.parse_args(["create", "o/r", "--repo-type", "model", "--visibility", vis]) assert args.visibility == vis def test_invalid_visibility_rejected(self, parser): with pytest.raises(SystemExit): - parser.parse_args( - ["create", "o/r", "--repo-type", "model", "--visibility", "secret"] - ) + parser.parse_args(["create", "o/r", "--repo-type", "model", "--visibility", "secret"]) def test_license_flag(self, parser): - args = parser.parse_args( - ["create", "o/r", "--repo-type", "model", "--license", "apache-2.0"] - ) + args = parser.parse_args(["create", "o/r", "--repo-type", "model", "--license", "apache-2.0"]) assert args.license == "apache-2.0" def test_chinese_name(self, parser): - args = parser.parse_args( - ["create", "o/r", "--repo-type", "model", "--chinese-name", "测试模型"] - ) + args = parser.parse_args(["create", "o/r", "--repo-type", "model", "--chinese-name", "测试模型"]) assert args.chinese_name == "测试模型" def test_chinese_name_underscore(self, parser): - args = parser.parse_args( - ["create", "o/r", "--repo-type", "model", "--chinese_name", "测试"] - ) + args = parser.parse_args(["create", "o/r", "--repo-type", "model", "--chinese_name", "测试"]) assert args.chinese_name == "测试" def test_description(self, parser): - args = parser.parse_args( - ["create", "o/r", "--repo-type", "model", "--description", "A test model"] - ) + args = parser.parse_args(["create", "o/r", "--repo-type", "model", "--description", "A test model"]) assert args.description == "A test model" def test_exist_ok_flag(self, parser): - args = parser.parse_args( - ["create", "o/r", "--repo-type", "model", "--exist-ok"] - ) + args = parser.parse_args(["create", "o/r", "--repo-type", "model", "--exist-ok"]) assert args.exist_ok is True def test_exist_ok_underscore(self, parser): - args = parser.parse_args( - ["create", "o/r", "--repo-type", "model", "--exist_ok"] - ) + args = parser.parse_args(["create", "o/r", "--repo-type", "model", "--exist_ok"]) assert args.exist_ok is True def test_exist_ok_default_false(self, parser): @@ -102,62 +87,72 @@ def test_missing_repo_type_exits(self, parser): @pytest.mark.parametrize("sdk", ["gradio", "streamlit", "docker", "static"]) def test_studio_sdk_type(self, parser, sdk): - args = parser.parse_args( - ["create", "o/r", "--repo-type", "studio", "--sdk-type", sdk] - ) + args = parser.parse_args(["create", "o/r", "--repo-type", "studio", "--sdk-type", sdk]) assert args.sdk_type == sdk def test_invalid_sdk_type_rejected(self, parser): with pytest.raises(SystemExit): - parser.parse_args( - ["create", "o/r", "--repo-type", "studio", "--sdk-type", "flask"] - ) + parser.parse_args(["create", "o/r", "--repo-type", "studio", "--sdk-type", "flask"]) def test_studio_sdk_version(self, parser): - args = parser.parse_args( - ["create", "o/r", "--repo-type", "studio", "--sdk-version", "4.0"] - ) + args = parser.parse_args(["create", "o/r", "--repo-type", "studio", "--sdk-version", "4.0"]) assert args.sdk_version == "4.0" def test_studio_base_image(self, parser): - args = parser.parse_args( - ["create", "o/r", "--repo-type", "studio", "--base-image", "python:3.11"] - ) + args = parser.parse_args(["create", "o/r", "--repo-type", "studio", "--base-image", "python:3.11"]) assert args.base_image == "python:3.11" def test_studio_cover_image(self, parser): - args = parser.parse_args( - ["create", "o/r", "--repo-type", "studio", "--cover-image", "https://img.png"] - ) + args = parser.parse_args(["create", "o/r", "--repo-type", "studio", "--cover-image", "https://img.png"]) assert args.cover_image == "https://img.png" def test_studio_hardware(self, parser): - args = parser.parse_args( - ["create", "o/r", "--repo-type", "studio", "--hardware", "gpu.a10"] - ) + args = parser.parse_args(["create", "o/r", "--repo-type", "studio", "--hardware", "gpu.a10"]) assert args.hardware == "gpu.a10" def test_all_studio_options_combined(self, parser): - args = parser.parse_args([ - "create", "o/studio1", "--repo-type", "studio", - "--visibility", "private", - "--sdk-type", "gradio", "--sdk-version", "4.0", - "--base-image", "python:3.11", - "--cover-image", "https://img.png", - "--hardware", "gpu.a10", - "--license", "mit", - "--description", "demo", - "--chinese-name", "演示", - ]) + args = parser.parse_args( + [ + "create", + "o/studio1", + "--repo-type", + "studio", + "--visibility", + "private", + "--sdk-type", + "gradio", + "--sdk-version", + "4.0", + "--base-image", + "python:3.11", + "--cover-image", + "https://img.png", + "--hardware", + "gpu.a10", + "--license", + "mit", + "--description", + "demo", + "--chinese-name", + "演示", + ] + ) assert args.repo_type == "studio" assert args.sdk_type == "gradio" assert args.sdk_version == "4.0" assert args.hardware == "gpu.a10" def test_subcmd_token_endpoint(self, parser): - args = parser.parse_args([ - "create", "o/r", "--repo-type", "model", "--token", "tk", - ]) + args = parser.parse_args( + [ + "create", + "o/r", + "--repo-type", + "model", + "--token", + "tk", + ] + ) assert args.subcmd_token == "tk" @@ -179,9 +174,18 @@ def test_repo_type_required(self, parser): parser.parse_args(["info", "o/r"]) def test_subcmd_token_endpoint(self, parser): - args = parser.parse_args([ - "info", "o/r", "--repo-type", "model", "--token", "tk", "--endpoint", "https://x.cn", - ]) + args = parser.parse_args( + [ + "info", + "o/r", + "--repo-type", + "model", + "--token", + "tk", + "--endpoint", + "https://x.cn", + ] + ) assert args.subcmd_token == "tk" assert args.subcmd_endpoint == "https://x.cn" @@ -225,13 +229,20 @@ def test_missing_repo_type_exits(self, parser): args = parser.parse_args(["list"]) assert args.repo_type is None from modelscope_hub.cli.repo import ListCommand + with pytest.raises(SystemExit): ListCommand(args).execute() def test_subcmd_token_endpoint(self, parser): - args = parser.parse_args([ - "list", "--repo-type", "model", "--token", "tk", - ]) + args = parser.parse_args( + [ + "list", + "--repo-type", + "model", + "--token", + "tk", + ] + ) assert args.subcmd_token == "tk" @@ -265,9 +276,16 @@ def test_yes_default_false(self, parser): assert args.yes is False def test_subcmd_token_endpoint(self, parser): - args = parser.parse_args([ - "delete", "o/r", "--repo-type", "model", "--token", "tk", - ]) + args = parser.parse_args( + [ + "delete", + "o/r", + "--repo-type", + "model", + "--token", + "tk", + ] + ) assert args.subcmd_token == "tk" @@ -279,26 +297,47 @@ class TestCreateExecute: """CreateCommand.execute() logic.""" def test_create_model(self, parser, mock_api, capsys): - args = parser.parse_args([ - "create", "owner/my-model", "--repo-type", "model", - "--visibility", "private", "--license", "apache-2.0", - ]) + args = parser.parse_args( + [ + "create", + "owner/my-model", + "--repo-type", + "model", + "--visibility", + "private", + "--license", + "apache-2.0", + ] + ) with patch("modelscope_hub.cli.repo.make_api", return_value=mock_api): CreateCommand(args).execute() mock_api.create_repo.assert_called_once_with( - "owner/my-model", "model", - visibility="private", license="apache-2.0", - chinese_name=None, description=None, gated_mode=None, + "owner/my-model", + "model", + visibility="private", + license="apache-2.0", + chinese_name=None, + description=None, + gated_mode=None, ) out = capsys.readouterr().out assert "Created" in out def test_create_studio_with_extras(self, parser, mock_api, capsys): - args = parser.parse_args([ - "create", "owner/demo", "--repo-type", "studio", - "--sdk-type", "gradio", "--sdk-version", "4.0", - "--hardware", "gpu.a10", - ]) + args = parser.parse_args( + [ + "create", + "owner/demo", + "--repo-type", + "studio", + "--sdk-type", + "gradio", + "--sdk-version", + "4.0", + "--hardware", + "gpu.a10", + ] + ) with patch("modelscope_hub.cli.repo.make_api", return_value=mock_api): CreateCommand(args).execute() call_kwargs = mock_api.create_repo.call_args @@ -308,9 +347,15 @@ def test_create_studio_with_extras(self, parser, mock_api, capsys): def test_exist_ok_swallows_exist_error(self, parser, mock_api, capsys): mock_api.create_repo.side_effect = Exception("Repository already exists") - args = parser.parse_args([ - "create", "owner/repo", "--repo-type", "model", "--exist-ok", - ]) + args = parser.parse_args( + [ + "create", + "owner/repo", + "--repo-type", + "model", + "--exist-ok", + ] + ) with patch("modelscope_hub.cli.repo.make_api", return_value=mock_api): CreateCommand(args).execute() out = capsys.readouterr().out @@ -318,27 +363,46 @@ def test_exist_ok_swallows_exist_error(self, parser, mock_api, capsys): def test_exist_ok_reraises_non_exist_error(self, parser, mock_api): mock_api.create_repo.side_effect = Exception("Permission denied") - args = parser.parse_args([ - "create", "owner/repo", "--repo-type", "model", "--exist-ok", - ]) + args = parser.parse_args( + [ + "create", + "owner/repo", + "--repo-type", + "model", + "--exist-ok", + ] + ) with patch("modelscope_hub.cli.repo.make_api", return_value=mock_api): with pytest.raises(Exception, match="Permission denied"): CreateCommand(args).execute() def test_create_without_exist_ok_raises(self, parser, mock_api): mock_api.create_repo.side_effect = Exception("Repository already exists") - args = parser.parse_args([ - "create", "owner/repo", "--repo-type", "model", - ]) + args = parser.parse_args( + [ + "create", + "owner/repo", + "--repo-type", + "model", + ] + ) with patch("modelscope_hub.cli.repo.make_api", return_value=mock_api): with pytest.raises(Exception, match="already exists"): CreateCommand(args).execute() def test_chinese_name_and_description_forwarded(self, parser, mock_api, capsys): - args = parser.parse_args([ - "create", "owner/repo", "--repo-type", "model", - "--chinese-name", "测试模型", "--description", "A test model", - ]) + args = parser.parse_args( + [ + "create", + "owner/repo", + "--repo-type", + "model", + "--chinese-name", + "测试模型", + "--description", + "A test model", + ] + ) with patch("modelscope_hub.cli.repo.make_api", return_value=mock_api): CreateCommand(args).execute() call_kwargs = mock_api.create_repo.call_args @@ -346,10 +410,18 @@ def test_chinese_name_and_description_forwarded(self, parser, mock_api, capsys): assert call_kwargs.kwargs["description"] == "A test model" def test_base_image_and_cover_image_forwarded(self, parser, mock_api, capsys): - args = parser.parse_args([ - "create", "owner/studio1", "--repo-type", "studio", - "--base-image", "python:3.11", "--cover-image", "https://img.png", - ]) + args = parser.parse_args( + [ + "create", + "owner/studio1", + "--repo-type", + "studio", + "--base-image", + "python:3.11", + "--cover-image", + "https://img.png", + ] + ) with patch("modelscope_hub.cli.repo.make_api", return_value=mock_api): CreateCommand(args).execute() call_kwargs = mock_api.create_repo.call_args @@ -358,27 +430,48 @@ def test_base_image_and_cover_image_forwarded(self, parser, mock_api, capsys): def test_create_dataset(self, parser, mock_api, capsys): """Create a dataset repo via CLI.""" - args = parser.parse_args([ - "create", "owner/my-dataset", "--repo-type", "dataset", - "--visibility", "private", "--license", "cc-by-4.0", - "--description", "Test dataset", - ]) + args = parser.parse_args( + [ + "create", + "owner/my-dataset", + "--repo-type", + "dataset", + "--visibility", + "private", + "--license", + "cc-by-4.0", + "--description", + "Test dataset", + ] + ) with patch("modelscope_hub.cli.repo.make_api", return_value=mock_api): CreateCommand(args).execute() mock_api.create_repo.assert_called_once_with( - "owner/my-dataset", "dataset", - visibility="private", license="cc-by-4.0", - chinese_name=None, description="Test dataset", gated_mode=None, + "owner/my-dataset", + "dataset", + visibility="private", + license="cc-by-4.0", + chinese_name=None, + description="Test dataset", + gated_mode=None, ) out = capsys.readouterr().out assert "Created" in out def test_create_dataset_public_with_license(self, parser, mock_api, capsys): """Create a public dataset with specific license.""" - args = parser.parse_args([ - "create", "owner/public-ds", "--repo-type", "dataset", - "--visibility", "public", "--license", "mit", - ]) + args = parser.parse_args( + [ + "create", + "owner/public-ds", + "--repo-type", + "dataset", + "--visibility", + "public", + "--license", + "mit", + ] + ) with patch("modelscope_hub.cli.repo.make_api", return_value=mock_api): CreateCommand(args).execute() call_kwargs = mock_api.create_repo.call_args @@ -388,32 +481,47 @@ def test_create_dataset_public_with_license(self, parser, mock_api, capsys): def test_create_dataset_with_chinese_name(self, parser, mock_api, capsys): """Create a dataset with chinese name and description.""" - args = parser.parse_args([ - "create", "owner/cn-dataset", "--repo-type", "dataset", - "--chinese-name", "测试数据集", - "--description", "这是一个测试数据集", - ]) + args = parser.parse_args( + [ + "create", + "owner/cn-dataset", + "--repo-type", + "dataset", + "--chinese-name", + "测试数据集", + "--description", + "这是一个测试数据集", + ] + ) with patch("modelscope_hub.cli.repo.make_api", return_value=mock_api): CreateCommand(args).execute() mock_api.create_repo.assert_called_once_with( - "owner/cn-dataset", "dataset", - visibility=None, license=None, - chinese_name="测试数据集", description="这是一个测试数据集", gated_mode=None, + "owner/cn-dataset", + "dataset", + visibility=None, + license=None, + chinese_name="测试数据集", + description="这是一个测试数据集", + gated_mode=None, ) - def test_create_skill_with_skill_file(self, parser, mock_api, capsys, tmp_path): """Skill file is uploaded and its ID is forwarded to create_repo.""" zip_file = tmp_path / "skill.zip" zip_file.write_bytes(b"PK dummy") - mock_api.upload_file_to_openapi = MagicMock( - return_value="8c378570-8991-431b-a82c-96f3d0b4f0f4" - ) - args = parser.parse_args([ - "create", "owner/my-skill", "--repo-type", "skill", - "--category", "developer-tools", - "--skill-file", str(zip_file), - ]) + mock_api.upload_file_to_openapi = MagicMock(return_value="8c378570-8991-431b-a82c-96f3d0b4f0f4") + args = parser.parse_args( + [ + "create", + "owner/my-skill", + "--repo-type", + "skill", + "--category", + "developer-tools", + "--skill-file", + str(zip_file), + ] + ) with patch("modelscope_hub.cli.repo.make_api", return_value=mock_api): CreateCommand(args).execute() mock_api.upload_file_to_openapi.assert_called_once() @@ -423,11 +531,18 @@ def test_create_skill_with_skill_file(self, parser, mock_api, capsys, tmp_path): def test_create_skill_file_not_found(self, parser, mock_api): """Non-existent --skill-file path causes SystemExit(2).""" - args = parser.parse_args([ - "create", "owner/my-skill", "--repo-type", "skill", - "--category", "developer-tools", - "--skill-file", "/nonexistent/skill.zip", - ]) + args = parser.parse_args( + [ + "create", + "owner/my-skill", + "--repo-type", + "skill", + "--category", + "developer-tools", + "--skill-file", + "/nonexistent/skill.zip", + ] + ) with patch("modelscope_hub.cli.repo.make_api", return_value=mock_api): with pytest.raises(SystemExit) as exc_info: CreateCommand(args).execute() @@ -453,20 +568,35 @@ class TestListExecute: """ListCommand.execute() logic.""" def test_list_with_results(self, parser, mock_api, capsys): - args = parser.parse_args([ - "list", "--repo-type", "model", "--owner", "org", "--page-size", "20", - ]) + args = parser.parse_args( + [ + "list", + "--repo-type", + "model", + "--owner", + "org", + "--page-size", + "20", + ] + ) with patch("modelscope_hub.cli.repo.make_api", return_value=mock_api): ListCommand(args).execute() mock_api.list_repos.assert_called_once_with( - "model", owner="org", search=None, page_number=1, page_size=20, + "model", + owner="org", + search=None, + page_number=1, + page_size=20, ) out = capsys.readouterr().out assert "owner/model1" in out def test_list_empty(self, parser, mock_api, capsys): mock_api.list_repos.return_value = PagedResult( - items=[], total_count=0, page_number=1, page_size=10, + items=[], + total_count=0, + page_number=1, + page_size=10, ) args = parser.parse_args(["list", "--repo-type", "model"]) with patch("modelscope_hub.cli.repo.make_api", return_value=mock_api): @@ -475,28 +605,49 @@ def test_list_empty(self, parser, mock_api, capsys): assert "no repositories found" in out def test_list_with_search(self, parser, mock_api, capsys): - args = parser.parse_args([ - "list", "--repo-type", "dataset", "--search", "qwen", - ]) + args = parser.parse_args( + [ + "list", + "--repo-type", + "dataset", + "--search", + "qwen", + ] + ) with patch("modelscope_hub.cli.repo.make_api", return_value=mock_api): ListCommand(args).execute() mock_api.list_repos.assert_called_once_with( - "dataset", owner=None, search="qwen", page_number=1, page_size=10, + "dataset", + owner=None, + search="qwen", + page_number=1, + page_size=10, ) def test_list_all_paginates(self, parser, mock_api, capsys): page1 = PagedResult( items=[RepoInfo(id=1, owner="o", name="m1", repo_type="model", downloads=10, likes=1)], - total_count=2, page_number=1, page_size=1, + total_count=2, + page_number=1, + page_size=1, ) page2 = PagedResult( items=[RepoInfo(id=2, owner="o", name="m2", repo_type="model", downloads=5, likes=0)], - total_count=2, page_number=2, page_size=1, + total_count=2, + page_number=2, + page_size=1, ) mock_api.list_repos.side_effect = [page1, page2] - args = parser.parse_args([ - "list", "--repo-type", "model", "--all", "--page-size", "1", - ]) + args = parser.parse_args( + [ + "list", + "--repo-type", + "model", + "--all", + "--page-size", + "1", + ] + ) with patch("modelscope_hub.cli.repo.make_api", return_value=mock_api): ListCommand(args).execute() assert mock_api.list_repos.call_count == 2 @@ -507,7 +658,10 @@ def test_list_all_paginates(self, parser, mock_api, capsys): def test_list_all_empty(self, parser, mock_api, capsys): mock_api.list_repos.return_value = PagedResult( - items=[], total_count=0, page_number=1, page_size=50, + items=[], + total_count=0, + page_number=1, + page_size=50, ) args = parser.parse_args(["list", "--repo-type", "model", "--all"]) with patch("modelscope_hub.cli.repo.make_api", return_value=mock_api): @@ -525,9 +679,15 @@ class TestDeleteExecute: """DeleteCommand.execute() logic.""" def test_delete_with_yes(self, parser, mock_api, capsys): - args = parser.parse_args([ - "delete", "owner/repo", "--repo-type", "model", "--yes", - ]) + args = parser.parse_args( + [ + "delete", + "owner/repo", + "--repo-type", + "model", + "--yes", + ] + ) with patch("modelscope_hub.cli.repo.make_api", return_value=mock_api): DeleteCommand(args).execute() mock_api.delete_repo.assert_called_once_with("owner/repo", "model") @@ -562,9 +722,16 @@ class TestRepoCompat: """Verify ``ms repo create/info/list/delete`` still works.""" def test_repo_create(self, parser): - args = parser.parse_args([ - "repo", "create", "o/r", "--repo-type", "model", "--exist-ok", - ]) + args = parser.parse_args( + [ + "repo", + "create", + "o/r", + "--repo-type", + "model", + "--exist-ok", + ] + ) assert args.repo_id == "o/r" assert args.exist_ok is True diff --git a/tests/cli/test_secret.py b/tests/cli/test_secret.py index 1200221..2f5b6cb 100644 --- a/tests/cli/test_secret.py +++ b/tests/cli/test_secret.py @@ -5,6 +5,7 @@ - Execution tests: mock HubApi for secret CRUD logic - Remote tests: real API lifecycle (existing) """ + from __future__ import annotations import warnings @@ -47,15 +48,31 @@ def test_missing_key_and_value_exits(self, parser): parser.parse_args(["secret", "add", "o/r"]) def test_explicit_repo_type(self, parser): - args = parser.parse_args([ - "secret", "add", "o/r", "K", "V", "--repo-type", "studio", - ]) + args = parser.parse_args( + [ + "secret", + "add", + "o/r", + "K", + "V", + "--repo-type", + "studio", + ] + ) assert args.repo_type == "studio" def test_subcmd_token_endpoint(self, parser): - args = parser.parse_args([ - "secret", "add", "o/r", "K", "V", "--token", "tk", - ]) + args = parser.parse_args( + [ + "secret", + "add", + "o/r", + "K", + "V", + "--token", + "tk", + ] + ) assert args.subcmd_token == "tk" diff --git a/tests/cli/test_upload.py b/tests/cli/test_upload.py index 9ff22fb..3d8f544 100644 --- a/tests/cli/test_upload.py +++ b/tests/cli/test_upload.py @@ -5,10 +5,10 @@ - Execution tests: mock HubApi for file/folder upload logic - Remote tests: real API upload (existing) """ + from __future__ import annotations import warnings -from pathlib import Path from unittest.mock import patch import pytest @@ -48,15 +48,27 @@ def test_invalid_repo_type_rejected(self, parser): parser.parse_args(["upload", "o/r", ".", "--repo-type", "studio"]) def test_commit_message(self, parser): - args = parser.parse_args([ - "upload", "o/r", ".", "--commit-message", "add weights", - ]) + args = parser.parse_args( + [ + "upload", + "o/r", + ".", + "--commit-message", + "add weights", + ] + ) assert args.commit_message == "add weights" def test_commit_description(self, parser): - args = parser.parse_args([ - "upload", "o/r", ".", "--commit-description", "extended desc", - ]) + args = parser.parse_args( + [ + "upload", + "o/r", + ".", + "--commit-description", + "extended desc", + ] + ) assert args.commit_description == "extended desc" def test_revision(self, parser): @@ -100,26 +112,45 @@ def test_disable_tqdm_default_false(self, parser): assert args.disable_tqdm is False def test_subcmd_token_endpoint(self, parser): - args = parser.parse_args([ - "upload", "o/r", ".", - "--token", "ms-tok", "--endpoint", "https://x.cn", - ]) + args = parser.parse_args( + [ + "upload", + "o/r", + ".", + "--token", + "ms-tok", + "--endpoint", + "https://x.cn", + ] + ) assert args.subcmd_token == "ms-tok" assert args.subcmd_endpoint == "https://x.cn" def test_all_options_combined(self, parser): - args = parser.parse_args([ - "upload", "my-org/my-model", "./output", "weights/", - "--repo-type", "dataset", - "--commit-message", "v2", - "--commit-description", "retrained", - "--revision", "dev", - "--include", "*.safetensors", - "--exclude", "*.ckpt", - "--max-workers", "4", - "--no-cache", - "--disable-tqdm", - ]) + args = parser.parse_args( + [ + "upload", + "my-org/my-model", + "./output", + "weights/", + "--repo-type", + "dataset", + "--commit-message", + "v2", + "--commit-description", + "retrained", + "--revision", + "dev", + "--include", + "*.safetensors", + "--exclude", + "*.ckpt", + "--max-workers", + "4", + "--no-cache", + "--disable-tqdm", + ] + ) assert args.repo_id == "my-org/my-model" assert args.local_path == "./output" assert args.path_in_repo == "weights/" @@ -189,11 +220,17 @@ def test_no_cache_forwarded(self, parser, mock_api, tmp_path, capsys): def test_commit_message_forwarded(self, parser, mock_api, tmp_path, capsys): test_file = tmp_path / "f.bin" test_file.write_text("data") - args = parser.parse_args([ - "upload", "o/r", str(test_file), - "--commit-message", "add weights", - "--commit-description", "desc", - ]) + args = parser.parse_args( + [ + "upload", + "o/r", + str(test_file), + "--commit-message", + "add weights", + "--commit-description", + "desc", + ] + ) with patch("modelscope_hub.cli.upload.make_api", return_value=mock_api): UploadCommand(args).execute() kw = mock_api.upload_file.call_args.kwargs @@ -203,9 +240,15 @@ def test_commit_message_forwarded(self, parser, mock_api, tmp_path, capsys): def test_dataset_upload(self, parser, mock_api, tmp_path, capsys): test_file = tmp_path / "data.csv" test_file.write_text("a,b\n1,2") - args = parser.parse_args([ - "upload", "o/r", str(test_file), "--repo-type", "dataset", - ]) + args = parser.parse_args( + [ + "upload", + "o/r", + str(test_file), + "--repo-type", + "dataset", + ] + ) with patch("modelscope_hub.cli.upload.make_api", return_value=mock_api): UploadCommand(args).execute() assert mock_api.upload_file.call_args.args[1] == "dataset" @@ -222,10 +265,17 @@ def test_include_exclude_forwarded(self, parser, mock_api, tmp_path, capsys): upload_dir = tmp_path / "src" upload_dir.mkdir() (upload_dir / "a.py").write_text("x") - args = parser.parse_args([ - "upload", "o/r", str(upload_dir), - "--include", "*.py", "--exclude", "*.pyc", - ]) + args = parser.parse_args( + [ + "upload", + "o/r", + str(upload_dir), + "--include", + "*.py", + "--exclude", + "*.pyc", + ] + ) with patch("modelscope_hub.cli.upload.make_api", return_value=mock_api): UploadCommand(args).execute() kw = mock_api.upload_folder.call_args.kwargs @@ -260,8 +310,7 @@ def test_disable_tqdm_forwarded(self, parser, mock_api, tmp_path, capsys): def test_resolve_paths_no_local_path(self, parser, mock_api, tmp_path, capsys): args = parser.parse_args(["upload", "o/nonexistent"]) cmd = UploadCommand(args) - with patch("os.path.isfile", return_value=False), \ - patch("os.path.isdir", return_value=False): + with patch("os.path.isfile", return_value=False), patch("os.path.isdir", return_value=False): local, pir = cmd._resolve_paths() assert local == "." assert pir is None diff --git a/tests/conftest.py b/tests/conftest.py index 4e2d041..374f0d1 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -1,4 +1,5 @@ """Global test configuration and fixtures.""" + from __future__ import annotations import os @@ -54,7 +55,9 @@ def is_remote_enabled() -> bool: # --------------------------------------------------------------------------- def pytest_configure(config): config.addinivalue_line("markers", "remote: tests requiring remote API access") - config.addinivalue_line("markers", "mock_only: tests using mock API (only run when MODELSCOPE_RUN_REMOTE_TESTS=false)") + config.addinivalue_line( + "markers", "mock_only: tests using mock API (only run when MODELSCOPE_RUN_REMOTE_TESTS=false)" + ) def pytest_collection_modifyitems(config, items): @@ -67,9 +70,7 @@ def pytest_collection_modifyitems(config, items): if remote_enabled: # Skip mock-only tests when real API is available - skip_mock = pytest.mark.skip( - reason="Mock-only tests skipped (remote mode active)" - ) + skip_mock = pytest.mark.skip(reason="Mock-only tests skipped (remote mode active)") for item in items: if "mock_only" in item.keywords: item.add_marker(skip_mock) diff --git a/tests/integration/conftest.py b/tests/integration/conftest.py index bf037f1..ed1d69e 100644 --- a/tests/integration/conftest.py +++ b/tests/integration/conftest.py @@ -1,4 +1,5 @@ """Integration test fixtures — real API calls with cleanup.""" + from __future__ import annotations import pytest diff --git a/tests/integration/run_all.py b/tests/integration/run_all.py index ee33072..a563380 100644 --- a/tests/integration/run_all.py +++ b/tests/integration/run_all.py @@ -18,6 +18,7 @@ test_remote_file_ops.py — File upload/download/delete with cleanup test_dataset_ops.py — Dataset-specific file listing and download """ + from __future__ import annotations import subprocess @@ -34,11 +35,14 @@ def main() -> int: quick = "--quick" in args cmd = [ - sys.executable, "-m", "pytest", + sys.executable, + "-m", + "pytest", str(_TESTS_DIR), "-v", "--tb=short", - "-m", "remote", + "-m", + "remote", ] if dry_run: diff --git a/tests/integration/test_dataset_ops.py b/tests/integration/test_dataset_ops.py index 0dae7ee..a951925 100644 --- a/tests/integration/test_dataset_ops.py +++ b/tests/integration/test_dataset_ops.py @@ -7,9 +7,8 @@ Uses a known small public dataset — no auth required for read operations. """ -from __future__ import annotations -from pathlib import Path +from __future__ import annotations import pytest @@ -70,10 +69,13 @@ def test_cli_download_dataset_single_file(self, test_endpoint, tmp_path): exit_code, out, err = run_cli( [ - "download", PUBLIC_DATASET_ID, + "download", + PUBLIC_DATASET_ID, PUBLIC_DATASET_SMALL_FILE, - "--repo-type", "dataset", - "--local-dir", str(tmp_path), + "--repo-type", + "dataset", + "--local-dir", + str(tmp_path), ], endpoint=test_endpoint, ) @@ -86,10 +88,14 @@ def test_cli_download_dataset_snapshot(self, test_endpoint, tmp_path): exit_code, out, err = run_cli( [ - "download", PUBLIC_DATASET_ID, - "--repo-type", "dataset", - "--local-dir", str(tmp_path), - "--include", "*.json", + "download", + PUBLIC_DATASET_ID, + "--repo-type", + "dataset", + "--local-dir", + str(tmp_path), + "--include", + "*.json", ], endpoint=test_endpoint, ) diff --git a/tests/integration/test_openapi.py b/tests/integration/test_openapi.py index faf3c05..a4c0c5c 100644 --- a/tests/integration/test_openapi.py +++ b/tests/integration/test_openapi.py @@ -6,6 +6,7 @@ Requires MODELSCOPE_TEST_TOKEN and MODELSCOPE_TEST_OWNER in tests/.env. """ + from __future__ import annotations import pytest diff --git a/tests/integration/test_remote_file_ops.py b/tests/integration/test_remote_file_ops.py index 7ae6565..282bd5e 100644 --- a/tests/integration/test_remote_file_ops.py +++ b/tests/integration/test_remote_file_ops.py @@ -3,6 +3,7 @@ These tests hit the real ModelScope API. They require MODELSCOPE_TEST_TOKEN and MODELSCOPE_TEST_OWNER in tests/.env. """ + from __future__ import annotations import pytest @@ -70,9 +71,7 @@ def test_delete_files(self): """delete_files removes the file from the repo.""" print(f"\n** repo_id: {self.repo_id}") print("** Deleting test_file.txt ...") - result = self.api.delete_files( - self.repo_id, "model", ["test_file.txt"], commit_message="cleanup" - ) + result = self.api.delete_files(self.repo_id, "model", ["test_file.txt"], commit_message="cleanup") print(f"** delete_files response: {result}") assert "test_file.txt" in result["deleted_files"] files = self.api.list_repo_files(self.repo_id, "model") diff --git a/tests/integration/test_remote_repo.py b/tests/integration/test_remote_repo.py index 89b1021..25ecc82 100644 --- a/tests/integration/test_remote_repo.py +++ b/tests/integration/test_remote_repo.py @@ -3,6 +3,7 @@ These tests create real repos on ModelScope and clean up after themselves. Requires MODELSCOPE_TEST_TOKEN and MODELSCOPE_TEST_OWNER in tests/.env. """ + from __future__ import annotations import pytest @@ -40,7 +41,7 @@ def test_create_and_delete_model_repo(self, api, test_owner, unique_repo_name): # Cleanup try: api.delete_repo(repo_id, "model") - print(f"** delete_repo: success") + print("** delete_repo: success") except Exception as e: print(f"** delete_repo: failed - {e}") @@ -67,7 +68,7 @@ def test_create_and_delete_dataset_repo(self, api, test_owner, unique_repo_name) # Cleanup try: api.delete_repo(repo_id, "dataset") - print(f"** delete_repo: success") + print("** delete_repo: success") except Exception as e: print(f"** delete_repo: failed - {e}") diff --git a/tests/integration/test_sdk_api.py b/tests/integration/test_sdk_api.py index f3ddc1e..769daef 100644 --- a/tests/integration/test_sdk_api.py +++ b/tests/integration/test_sdk_api.py @@ -6,15 +6,15 @@ Requires MODELSCOPE_TEST_TOKEN and MODELSCOPE_TEST_OWNER in tests/.env. """ + from __future__ import annotations -import tempfile import warnings from pathlib import Path import pytest -from modelscope_hub import HubApi, RepoType +from modelscope_hub import HubApi from modelscope_hub.errors import NotExistError @@ -49,7 +49,8 @@ def test_create_get_delete_dataset(self, api, test_owner, unique_repo_name): repo_id = f"{test_owner}/{unique_repo_name}_ds" try: info = api.create_repo( - repo_id, "dataset", + repo_id, + "dataset", visibility="private", license="cc-by-4.0", ) @@ -96,28 +97,36 @@ def setup_repo(self, api, test_owner, unique_repo_name): def test_upload_and_download_file(self, tmp_path): self.api.upload_file( - self.repo_id, "model", + self.repo_id, + "model", b"test content for sdk", "sdk_test.txt", commit_message="sdk test upload", ) local = self.api.download_file( - self.repo_id, "model", "sdk_test.txt", - cache_dir=str(tmp_path), force=True, + self.repo_id, + "model", + "sdk_test.txt", + cache_dir=str(tmp_path), + force=True, ) assert local.exists() assert local.read_text() == "test content for sdk" def test_download_file_to_local_dir(self, tmp_path): self.api.upload_file( - self.repo_id, "model", + self.repo_id, + "model", b"local dir content", "subdir/data.txt", commit_message="upload for local_dir test", ) local = self.api.download_file( - self.repo_id, "model", "subdir/data.txt", - local_dir=str(tmp_path), force=True, + self.repo_id, + "model", + "subdir/data.txt", + local_dir=str(tmp_path), + force=True, ) expected = tmp_path / "subdir" / "data.txt" assert local == expected @@ -126,13 +135,23 @@ def test_download_file_to_local_dir(self, tmp_path): def test_download_repo_snapshot(self, tmp_path): self.api.upload_file( - self.repo_id, "model", b"file1", "a.txt", commit_message="a", + self.repo_id, + "model", + b"file1", + "a.txt", + commit_message="a", ) self.api.upload_file( - self.repo_id, "model", b"file2", "b.txt", commit_message="b", + self.repo_id, + "model", + b"file2", + "b.txt", + commit_message="b", ) output = self.api.download_repo( - self.repo_id, "model", cache_dir=str(tmp_path), + self.repo_id, + "model", + cache_dir=str(tmp_path), ) assert output.is_dir() files = [p.name for p in output.rglob("*") if p.is_file()] @@ -141,23 +160,38 @@ def test_download_repo_snapshot(self, tmp_path): def test_download_repo_to_local_dir(self, tmp_path): self.api.upload_file( - self.repo_id, "model", b"x", "x.txt", commit_message="x", + self.repo_id, + "model", + b"x", + "x.txt", + commit_message="x", ) output = self.api.download_repo( - self.repo_id, "model", local_dir=str(tmp_path / "out"), + self.repo_id, + "model", + local_dir=str(tmp_path / "out"), ) assert output == tmp_path / "out" assert (tmp_path / "out" / "x.txt").exists() def test_download_repo_with_patterns(self, tmp_path): self.api.upload_file( - self.repo_id, "model", b"bin", "weights.bin", commit_message="bin", + self.repo_id, + "model", + b"bin", + "weights.bin", + commit_message="bin", ) self.api.upload_file( - self.repo_id, "model", b"json", "config.json", commit_message="json", + self.repo_id, + "model", + b"json", + "config.json", + commit_message="json", ) output = self.api.download_repo( - self.repo_id, "model", + self.repo_id, + "model", cache_dir=str(tmp_path), allow_patterns=["*.json"], ) @@ -167,7 +201,11 @@ def test_download_repo_with_patterns(self, tmp_path): def test_list_repo_files(self): self.api.upload_file( - self.repo_id, "model", b"data", "list_test.txt", commit_message="list", + self.repo_id, + "model", + b"data", + "list_test.txt", + commit_message="list", ) files = self.api.list_repo_files(self.repo_id, "model") paths = [f.path for f in files] @@ -176,10 +214,17 @@ def test_list_repo_files(self): @pytest.mark.xfail(reason="Server restricts file deletion to cookie-based session auth") def test_delete_files(self): self.api.upload_file( - self.repo_id, "model", b"del", "to_delete.txt", commit_message="del", + self.repo_id, + "model", + b"del", + "to_delete.txt", + commit_message="del", ) self.api.delete_files( - self.repo_id, "model", ["to_delete.txt"], commit_message="cleanup", + self.repo_id, + "model", + ["to_delete.txt"], + commit_message="cleanup", ) files = self.api.list_repo_files(self.repo_id, "model") paths = [f.path for f in files] @@ -192,7 +237,8 @@ def test_upload_folder(self, tmp_path): (folder / "f2.txt").write_text("two") self.api.upload_folder( - self.repo_id, "model", + self.repo_id, + "model", str(folder), path_in_repo="", commit_message="folder upload", @@ -213,7 +259,11 @@ def setup_repo(self, api, test_owner, unique_repo_name): self.api = api api.create_repo(self.repo_id, "model", visibility="private") api.upload_file( - self.repo_id, "model", b"init", "init.txt", commit_message="initial", + self.repo_id, + "model", + b"init", + "init.txt", + commit_message="initial", ) yield with warnings.catch_warnings(): @@ -288,9 +338,11 @@ def test_scan_cache_on_tmp(self, tmp_path): def test_download_then_scan_cache(self, api, tmp_path): api.download_file( - "Qwen/Qwen2.5-0.5B", "model", + "Qwen/Qwen2.5-0.5B", + "model", "config.json", - cache_dir=str(tmp_path), force=True, + cache_dir=str(tmp_path), + force=True, ) report = api.scan_cache(cache_dir=str(tmp_path)) assert report.total_repos >= 1 @@ -298,9 +350,11 @@ def test_download_then_scan_cache(self, api, tmp_path): def test_clear_cache_by_type(self, api, tmp_path): api.download_file( - "Qwen/Qwen2.5-0.5B", "model", + "Qwen/Qwen2.5-0.5B", + "model", "config.json", - cache_dir=str(tmp_path), force=True, + cache_dir=str(tmp_path), + force=True, ) freed = api.clear_cache(cache_dir=str(tmp_path), repo_type="model") assert freed >= 0 @@ -387,7 +441,6 @@ def test_get_model_branches_and_tags(self, test_token, test_endpoint): def test_get_valid_revision_nonexistent_raises(self, test_token, test_endpoint): from modelscope_hub.compat import LegacyHubApi - from modelscope_hub.errors import NotExistError legacy = LegacyHubApi(token=test_token, endpoint=test_endpoint) with pytest.raises(NotExistError): diff --git a/tests/test_compat_get_model_files.py b/tests/test_compat_get_model_files.py index 1ab39dc..e6c098c 100644 --- a/tests/test_compat_get_model_files.py +++ b/tests/test_compat_get_model_files.py @@ -5,6 +5,7 @@ guard for callers (e.g. vLLM) that pass the historical ``revision`` / ``root`` keyword arguments. """ + from __future__ import annotations from types import SimpleNamespace @@ -24,42 +25,35 @@ def _fake_files(): class TestGetModelFilesLegacyCompat: def test_revision_is_accepted_and_forwarded(self): lha = LegacyHubApi() - with mock.patch.object( - lha._api, "list_repo_files", - return_value=_fake_files()) as m: - out = lha.get_model_files( - "Qwen/Qwen2.5-1.5B-Instruct", revision="v2") + with mock.patch.object(lha._api, "list_repo_files", return_value=_fake_files()) as m: + out = lha.get_model_files("Qwen/Qwen2.5-1.5B-Instruct", revision="v2") assert [f["Path"] for f in out] == [ - "config.json", "model.safetensors", "subdir/extra.bin", + "config.json", + "model.safetensors", + "subdir/extra.bin", ] _, kwargs = m.call_args assert kwargs["revision"] == "v2" def test_root_restricts_to_subpath(self): lha = LegacyHubApi() - with mock.patch.object( - lha._api, "list_repo_files", return_value=_fake_files()): + with mock.patch.object(lha._api, "list_repo_files", return_value=_fake_files()): out = lha.get_model_files("owner/name", root="subdir") assert [f["Path"] for f in out] == ["subdir/extra.bin"] def test_tolerates_legacy_transport_kwargs(self): lha = LegacyHubApi() - with mock.patch.object( - lha._api, "list_repo_files", return_value=_fake_files()): + with mock.patch.object(lha._api, "list_repo_files", return_value=_fake_files()): # Historical kwargs must not raise "unexpected keyword argument". - out = lha.get_model_files( - "owner/name", revision="master", - use_cookies=True, headers={}) + out = lha.get_model_files("owner/name", revision="master", use_cookies=True, headers={}) assert len(out) == 3 def test_default_revision_none_forwarded(self): lha = LegacyHubApi() - with mock.patch.object( - lha._api, "list_repo_files", - return_value=_fake_files()) as m: + with mock.patch.object(lha._api, "list_repo_files", return_value=_fake_files()) as m: lha.get_model_files("owner/name") _, kwargs = m.call_args diff --git a/tests/test_compat_snapshot_download.py b/tests/test_compat_snapshot_download.py index 389f1ba..8f5b85b 100644 --- a/tests/test_compat_snapshot_download.py +++ b/tests/test_compat_snapshot_download.py @@ -6,6 +6,7 @@ whole ``compat -> HubApi facade -> DownloadManager`` forwarding chain is exercised (a facade that drops ``progress_callbacks`` would fail here). """ + from __future__ import annotations from unittest import mock @@ -21,9 +22,7 @@ class _DummyCallback(ProgressCallback): class TestSnapshotDownloadProgressCallbacks: def test_progress_callbacks_forwarded_through_facade(self): - with mock.patch.object( - DownloadManager, "download_repo", - return_value="/tmp/snapshot") as m: + with mock.patch.object(DownloadManager, "download_repo", return_value="/tmp/snapshot") as m: result = snapshot_download( "owner/repo", progress_callbacks=[_DummyCallback], @@ -35,9 +34,7 @@ def test_progress_callbacks_forwarded_through_facade(self): assert kwargs["progress_callbacks"] == [_DummyCallback] def test_progress_callbacks_default_none(self): - with mock.patch.object( - DownloadManager, "download_repo", - return_value="/tmp/snapshot") as m: + with mock.patch.object(DownloadManager, "download_repo", return_value="/tmp/snapshot") as m: snapshot_download("owner/repo", local_files_only=True) _, kwargs = m.call_args diff --git a/tests/test_config_token.py b/tests/test_config_token.py index 49521e6..f59d0cf 100644 --- a/tests/test_config_token.py +++ b/tests/test_config_token.py @@ -8,6 +8,7 @@ Regression test for the bug where ``MODELSCOPE_API_TOKEN="" ms-hub agent upload ...`` uploaded successfully by silently reusing the stored credential. """ + from __future__ import annotations from unittest.mock import patch diff --git a/tests/test_legacy_cache_detection.py b/tests/test_legacy_cache_detection.py index e7c46f3..e76d46c 100644 --- a/tests/test_legacy_cache_detection.py +++ b/tests/test_legacy_cache_detection.py @@ -4,6 +4,7 @@ which lets ``download_repo`` / ``download_file`` reuse an existing old-SDK cache instead of re-downloading into the new layout. """ + from __future__ import annotations from modelscope_hub.api import HubApi @@ -46,13 +47,11 @@ def test_multi_dot_name_encoding(self, tmp_path): def test_clean_cache_returns_none(self, tmp_path): dm = _make_download_manager() - assert dm._find_legacy_repo_dir( - "Qwen/Qwen3.5-4B", "model", tmp_path) is None + assert dm._find_legacy_repo_dir("Qwen/Qwen3.5-4B", "model", tmp_path) is None def test_empty_legacy_dir_returns_none(self, tmp_path): legacy = tmp_path / "models" / "Qwen" / "Qwen3___5-4B" legacy.mkdir(parents=True) # exists but empty dm = _make_download_manager() - assert dm._find_legacy_repo_dir( - "Qwen/Qwen3.5-4B", "model", tmp_path) is None + assert dm._find_legacy_repo_dir("Qwen/Qwen3.5-4B", "model", tmp_path) is None diff --git a/tests/test_upload_lfs_gate.py b/tests/test_upload_lfs_gate.py index 976e3eb..225790e 100644 --- a/tests/test_upload_lfs_gate.py +++ b/tests/test_upload_lfs_gate.py @@ -95,16 +95,18 @@ def test_upload_folder_cached_normal_hash_skips_batch_blob_validation(tmp_path: cache_key = f"README.md|{st.st_mtime}|{st.st_size}" cache_path = tmp_path / ".ms_upload_cache" cache_path.write_text( - json.dumps({ - "version": 3, - "repo_id": "owner/repo", - "files": { - cache_key: { - "hash": hashlib.sha256(content).hexdigest(), - "size": st.st_size, + json.dumps( + { + "version": 3, + "repo_id": "owner/repo", + "files": { + cache_key: { + "hash": hashlib.sha256(content).hexdigest(), + "size": st.st_size, + }, }, - }, - }), + } + ), encoding="utf-8", ) diff --git a/tests/test_utils.py b/tests/test_utils.py index 008a3b5..b6fac64 100644 --- a/tests/test_utils.py +++ b/tests/test_utils.py @@ -191,10 +191,19 @@ def test_naive_string(self): sh_tz = zoneinfo.ZoneInfo("Asia/Shanghai") assert dt.tzinfo == sh_tz - def test_datetime_passthrough(self): + def test_aware_datetime_normalized(self): + # parse_timestamp normalizes aware datetimes to the target timezone + # (astimezone always returns a new object): same instant, new tzinfo. original = datetime(2024, 6, 1, 12, 0, tzinfo=timezone.utc) result = parse_timestamp(original) - assert result is original + assert result == original + assert result.tzinfo == zoneinfo.ZoneInfo("Asia/Shanghai") + + def test_naive_datetime_gets_target_tz(self): + naive = datetime(2024, 6, 1, 12, 0) + result = parse_timestamp(naive) + assert result.tzinfo == zoneinfo.ZoneInfo("Asia/Shanghai") + assert (result.year, result.hour) == (2024, 12) def test_invalid_string(self): with pytest.raises(ValueError): From 402f7c55c58dcebe4670d1eda9ef81f3cdb67686 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E7=8F=AD=E6=89=AC?= Date: Fri, 31 Jul 2026 17:43:30 +0800 Subject: [PATCH 10/13] update readme --- README.md | 12 +++++++++--- 1 file changed, 9 insertions(+), 3 deletions(-) diff --git a/README.md b/README.md index 436278e..e33120c 100644 --- a/README.md +++ b/README.md @@ -33,11 +33,20 @@ The official Python SDK & CLI for [ModelScope Hub](https://modelscope.cn) — do ## News +**v0.1.9** (2026-07-31) +- **Fix**: two stale unit tests that failed in downstream distro sandboxes ([#46](https://github.com/modelscope/modelscope_hub/issues/46), NixOS): align the `mcp deploy` call-shape assertion and the `parse_timestamp` timezone-normalization contract +- **Feature**: `ms-hub mcp deploy` prints the operational URL after deployment; new `--auth-check` and repeatable `--env KEY=VALUE` options; `--transport-type` now validates against `sse`/`streamable_http` +- **CI**: new `citest` workflow — mock-mode test suite (no credentials/network, mirrors distro build sandboxes) on Python 3.10/3.12/3.14 plus ruff/mypy hard gates; releases now require a green test gate before publishing; added a `pre-commit` config +- **Quality**: ruff & mypy debt cleared to zero; lint/type-check targets aligned to the supported Python floor (3.10) + **v0.1.8** (2026-07-21) - **Feature**: `ms-hub agent` raw file transfer (download/upload/list) for remote agent repos; visibility support for agent hub; cache checksum verification (`ms-hub cache verify`) - **Fix**: forward `progress_callbacks` through `HubApi.download_repo` so custom download-progress callbacks work end-to-end; harden legacy (pre-1.38) cache auto-detection (reuse existing `{cache}/models/...` and default `{cache}/hub/models/...` layouts); normal (non-LFS) file upload - **Packaging**: rename console scripts to `modelscope-hub` / `ms-hub` to avoid a file conflict with the `modelscope` package (e.g. FreeBSD pkg) +
+Older releases + **v0.1.7** (2026-07-07) - **Feature**: intra-/inter-region cloud download acceleration, with a source marker in the progress bar - **Fix**: align `snapshot_download` cache path with the CLI; add legacy cache fallback @@ -49,9 +58,6 @@ The official Python SDK & CLI for [ModelScope Hub](https://modelscope.cn) — do **v0.1.5** (2026-06-30) - **Fix**: adaptive commit batch size for uploads -
-Older releases - **v0.1.4** (2026-06-26) - **Feature**: `gated_mode` parameter for `create_repo`; `ms-hub create --gated/--no-gated` flags - **Refactor**: unify visibility / gated_mode semantics in the SDK layer From 2375141a972180d1dcb533da35ff8e20eaf73528 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E7=8F=AD=E6=89=AC?= Date: Fri, 31 Jul 2026 18:51:11 +0800 Subject: [PATCH 11/13] fix 3.10 citest --- tests/cli/test_compat_cache_dir.py | 16 ++++++++++++---- 1 file changed, 12 insertions(+), 4 deletions(-) diff --git a/tests/cli/test_compat_cache_dir.py b/tests/cli/test_compat_cache_dir.py index ed45f47..429d1ae 100644 --- a/tests/cli/test_compat_cache_dir.py +++ b/tests/cli/test_compat_cache_dir.py @@ -8,11 +8,19 @@ from __future__ import annotations +import importlib from pathlib import Path from unittest.mock import MagicMock, patch import pytest +# ``modelscope_hub.compat.snapshot_download`` (the submodule) is shadowed by +# the same-named function re-exported in ``compat/__init__``. String targets +# like ``patch("modelscope_hub.compat.snapshot_download.HubApi")`` resolve to +# the *function* on Python 3.10 (mock walks attributes before importing +# submodules), so grab the real module object and use ``patch.object``. +_snapshot_download_mod = importlib.import_module("modelscope_hub.compat.snapshot_download") + # --------------------------------------------------------------------------- # _resolve_legacy_paths unit tests (pure logic, mock api) @@ -74,7 +82,7 @@ def test_cache_dir_trailing_slash_preserved(self): class TestSnapshotDownloadCacheCompat: """Verify path conversion propagates correctly to download_repo.""" - @patch("modelscope_hub.compat.snapshot_download.HubApi") + @patch.object(_snapshot_download_mod, "HubApi") def test_cache_dir_passed_through(self, MockHubApi): """cache_dir without local_dir -> download_repo gets cache_dir directly.""" from modelscope_hub.compat.snapshot_download import snapshot_download @@ -91,7 +99,7 @@ def test_cache_dir_passed_through(self, MockHubApi): assert call_kwargs["cache_dir"] == "/tmp/cache" assert call_kwargs["local_dir"] is None - @patch("modelscope_hub.compat.snapshot_download.HubApi") + @patch.object(_snapshot_download_mod, "HubApi") def test_local_dir_explicit_not_overridden(self, MockHubApi): """Explicit local_dir is passed through without modification.""" from modelscope_hub.compat.snapshot_download import snapshot_download @@ -107,7 +115,7 @@ def test_local_dir_explicit_not_overridden(self, MockHubApi): call_kwargs = mock_api.download_repo.call_args[1] assert call_kwargs["local_dir"] == "/custom/dir" - @patch("modelscope_hub.compat.snapshot_download.HubApi") + @patch.object(_snapshot_download_mod, "HubApi") def test_dataset_snapshot_download_cache_dir_passthrough(self, MockHubApi): """dataset_snapshot_download passes cache_dir through.""" from modelscope_hub.compat.snapshot_download import dataset_snapshot_download @@ -218,7 +226,7 @@ def test_no_args_uses_standard_cache(self, MockHubApi, tmp_path): assert call_kwargs["cache_dir"] is None assert call_kwargs["local_dir"] is None - @patch("modelscope_hub.compat.snapshot_download.HubApi") + @patch.object(_snapshot_download_mod, "HubApi") def test_snapshot_no_args_uses_standard_cache(self, MockHubApi): """snapshot_download with no explicit dirs -> standard cache layout.""" from modelscope_hub.compat.snapshot_download import snapshot_download From 2d082bd040ddcf3e8dc69f3be0249517afd2b13e Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E7=8F=AD=E6=89=AC?= Date: Sat, 1 Aug 2026 21:39:37 +0800 Subject: [PATCH 12/13] fix(auth): stop misreporting login failures and revoking credentials `HubApi.login` caught `HubError` -- the base class of every SDK exception -- and re-raised everything as `AuthenticationError("...token was rejected by the server")` while calling `clear_token()`. Three problems followed: * network, timeout and 5xx failures were all reported as an invalid token * the server's own message and RequestId were discarded, leaving no way to diagnose the failure or trace it server-side * a failed attempt deleted the credential already on disk, so a single mistyped token logged the user out of a working session Both ModelScope sites answer an unknown token with the same HTTP 400 and business code 10010103009, so the server cannot distinguish "invalid token" from "token issued by the other site" -- only the client knows which site it addressed, which is why that disambiguation now happens here. Changes: * classify 4xx by the server business code, mapping 10010103009 to AuthenticationError (E3001); 5xx stays a retryable ServerError so a transient outage is never reclassified as permanent * propagate non-authentication failures unchanged, preserving their type, error code, message and RequestId * leave persisted credentials untouched on failure and roll back only the in-memory state of the instance * probe the peer site when the endpoint was not pinned, and point at `--endpoint` when the token turns out to be valid there * centralise endpoint normalisation in `HubConfig.normalize_endpoint` so a bare or upper-case scheme no longer fails deep inside the transport layer * remove every credential artefact together on teardown, and propagate a cleared token to an already constructed legacy client * print the full error cause chain under `--verbose` Adds 29 regression tests that stub only the `requests` transport, so the config, facade, legacy client and error-translation layers all take part. --- src/modelscope_hub/api.py | 120 ++++++++++++-- src/modelscope_hub/cli/main.py | 54 +++++-- src/modelscope_hub/config.py | 49 ++++-- src/modelscope_hub/errors.py | 72 ++++++--- tests/cli/test_main.py | 78 ++++++++- tests/test_credential_lifecycle.py | 112 +++++++++++++ tests/test_login_failure_paths.py | 251 +++++++++++++++++++++++++++++ 7 files changed, 687 insertions(+), 49 deletions(-) create mode 100644 tests/test_credential_lifecycle.py create mode 100644 tests/test_login_failure_paths.py diff --git a/src/modelscope_hub/api.py b/src/modelscope_hub/api.py index 3ce9a35..071d449 100644 --- a/src/modelscope_hub/api.py +++ b/src/modelscope_hub/api.py @@ -156,7 +156,7 @@ def __init__( base._endpoint_overridden = was_overridden self._config = base if endpoint is not None: - self._config.endpoint = endpoint.rstrip("/") + self._config.endpoint = HubConfig.normalize_endpoint(endpoint) self._config._endpoint_overridden = True if token is not None: self._config.token = token @@ -188,7 +188,9 @@ def legacy(self) -> LegacyClient: endpoint=self._config.endpoint or DEFAULT_ENDPOINT, user_agent=build_user_agent(self._config.get_session_id()), ) - elif self._legacy.token != self._config.token and self._config.token: + elif self._legacy.token != self._config.token: + # Clears propagate as well as changes: a cached client left holding a + # revoked token would keep authenticating with it. self._legacy.token = self._config.token return self._legacy @@ -483,8 +485,19 @@ def login(self, token: str) -> UserInfo: InvalidParameter When ``token`` is empty or whitespace-only. AuthenticationError - When the server rejects the token. The bad token is cleared - from local storage before re-raising. + When the server rejects the token. The server's own explanation is + preserved, and an endpoint hint is appended when the token turns + out to be valid on the peer ModelScope site. + HubError + Transport, timeout and server-side failures propagate unchanged -- + they are never reported as a rejected token. + + Notes + ----- + A failed attempt leaves persisted credentials untouched. Until the + server has accepted the new token, the stored credential is still the + caller's only working one, so revoking it on failure would turn a + mistyped token into an unintended logout. Examples -------- @@ -497,6 +510,8 @@ def login(self, token: str) -> UserInfo: raise InvalidParameter("token must be a non-empty string") token = token.strip() + previous_token = self._config.token + previous_logged_out = self._config._logged_out self._config.token = token self._config._logged_out = False self._openapi = None @@ -505,12 +520,12 @@ def login(self, token: str) -> UserInfo: try: data, cookies = self.legacy.login(token) - except (AuthenticationError, HubError) as exc: - self._config.clear_token() - raise AuthenticationError( - "Login failed: the provided token was rejected by the server.", - status_code=getattr(exc, "status_code", None), - ) from exc + except HubError as exc: + self._restore_credential_state(previous_token, previous_logged_out) + explained = self._explain_login_failure(token, exc) + if explained is exc: + raise + raise explained from exc git_token = data.get("AccessToken", "") username = data.get("Username", "") @@ -524,6 +539,91 @@ def login(self, token: str) -> UserInfo: return self.whoami() + def _restore_credential_state(self, token: str | None, logged_out: bool) -> None: + """Roll the in-memory credential back to its pre-login value. + + Persisted credentials are deliberately left alone; only this instance's + transient state is rewound, so a failed attempt leaves the object + exactly as it was found instead of poisoning it with a rejected token. + """ + self._config.token = token + self._config._logged_out = logged_out + self._openapi = None + if self._legacy is not None: + self._legacy.token = token + + def _explain_login_failure(self, token: str, exc: HubError) -> HubError: + """Return the exception to surface for a failed login attempt. + + Only authentication failures are re-worded. Network, timeout and + server-side errors are handed back untouched, because presenting them + as a rejected token would send the caller after the wrong remedy. + + The two ModelScope sites keep separate account systems and answer an + unknown token with the same business code, so the server cannot tell + "invalid token" apart from "token issued by the other site". Only the + client knows which site it addressed, which is why that disambiguation + has to happen here. + """ + if not isinstance(exc, AuthenticationError): + return exc + peer = self._peer_site_endpoint() + if peer is None or not self._token_valid_on(token, peer): + return exc + return AuthenticationError( + f"{exc.message} This token is valid on {peer} instead; retry with " + f"--endpoint {peer} (or set MODELSCOPE_ENDPOINT={peer}).", + status_code=exc.status_code, + request_id=exc.request_id, + response_body=exc.response_body, + url=exc.url, + method=exc.method, + ) + + def _peer_site_endpoint(self) -> str | None: + """Return the sibling ModelScope site, or ``None`` when not applicable. + + An explicitly configured endpoint is always respected, mirroring + :meth:`resolve_endpoint_for_read`: when the caller has pinned a site we + do not second-guess it. + """ + if self._config._endpoint_overridden: + return None + from .constants import DEFAULT_INTL_ENDPOINT + + def site_key(url: str) -> str: + host = (urlparse(url).hostname or "").lower() + return host[4:] if host.startswith("www.") else host + + current = site_key(self._config.endpoint or DEFAULT_ENDPOINT) + for candidate in (DEFAULT_ENDPOINT, DEFAULT_INTL_ENDPOINT): + if site_key(candidate) != current: + return candidate + return None + + @staticmethod + def _token_valid_on(token: str, endpoint: str) -> bool: + """Best-effort check of whether *token* authenticates against *endpoint*. + + Runs on the failure path only and is strictly advisory: any error means + "cannot confirm", so a probe outage degrades to the plain server message + rather than producing a misleading hint. Retries are disabled to keep + the failure path responsive. + """ + from .constants import API_CONNECT_TIMEOUT + + probe = LegacyClient( + token=None, + endpoint=endpoint, + timeout=API_CONNECT_TIMEOUT, + max_retries=0, + ) + try: + probe.login(token) + except Exception: # advisory only -- never mask the original failure + return False + return True + def logout(self) -> None: """Clear the locally persisted token. diff --git a/src/modelscope_hub/cli/main.py b/src/modelscope_hub/cli/main.py index 6fa219d..1bd198e 100644 --- a/src/modelscope_hub/cli/main.py +++ b/src/modelscope_hub/cli/main.py @@ -89,7 +89,7 @@ def _build_parser() -> argparse.ArgumentParser: "-v", "--verbose", action="store_true", - help="Enable verbose (DEBUG) logging.", + help="Enable DEBUG logging and print the full error cause chain.", ) subparsers = parser.add_subparsers(dest="command", metavar="COMMAND") @@ -191,6 +191,45 @@ def _discover_plugins(subparsers) -> None: logging.getLogger(__name__).debug("Failed to load CLI plugin %r: %s", ep.name, exc) +# --------------------------------------------------------------------------- +# Error reporting +# --------------------------------------------------------------------------- +def _next_cause(exc: BaseException) -> BaseException | None: + """Return what *exc* was raised from, honouring ``raise ... from None``.""" + if exc.__cause__ is not None: + return exc.__cause__ + if exc.__suppress_context__: + return None + return exc.__context__ + + +def _report_hub_error(exc: HubError, *, verbose: bool, max_depth: int = 5) -> None: + """Print a structured report for an SDK error. + + ``str(exc)`` already carries the error code, HTTP status, request id and -- + for API errors -- the request/response detail. Verbose mode additionally + unwinds the cause chain: wrapping an exception is convenient for callers but + otherwise hides the originating failure from whoever has to diagnose it. + + The walk is bounded by *max_depth* and skips exceptions already visited, so + a self-referential chain cannot stall the error path. + """ + error(str(exc)) + if exc.suggestion and exc.error_code != "E9001": + info(f"Suggestion: {exc.suggestion}") + if not verbose: + return + + seen = {id(exc)} + cause = _next_cause(exc) + depth = 1 + while cause is not None and id(cause) not in seen and depth <= max_depth: + info(f"{' ' * depth}Caused by: {cause.__class__.__name__}: {cause}") + seen.add(id(cause)) + cause = _next_cause(cause) + depth += 1 + + # --------------------------------------------------------------------------- # Entry point # --------------------------------------------------------------------------- @@ -200,8 +239,9 @@ def run_cmd(argv: Sequence[str] | None = None) -> int: parser = _build_parser() args = parser.parse_args(argv) + verbose = bool(getattr(args, "verbose", False)) logging.basicConfig( - level=logging.DEBUG if getattr(args, "verbose", False) else logging.INFO, + level=logging.DEBUG if verbose else logging.INFO, format="%(levelname)s %(name)s: %(message)s", ) @@ -218,14 +258,10 @@ def run_cmd(argv: Sequence[str] | None = None) -> int: except SystemExit as exc: # honour explicit SystemExit from subcommands return int(exc.code) if isinstance(exc.code, int) else (0 if exc.code is None else 1) except (InvalidParameter, NotSupportedError) as exc: - error(str(exc)) - if exc.suggestion: - info(f"Suggestion: {exc.suggestion}") + _report_hub_error(exc, verbose=verbose) return 2 except HubError as exc: - error(str(exc)) - if exc.suggestion and exc.error_code != "E9001": - info(f"Suggestion: {exc.suggestion}") + _report_hub_error(exc, verbose=verbose) return 1 except ValueError as exc: error(str(exc)) @@ -235,7 +271,7 @@ def run_cmd(argv: Sequence[str] | None = None) -> int: return 2 except Exception as exc: # pragma: no cover - unexpected error(f"Unexpected error: {exc.__class__.__name__}: {exc}") - if getattr(args, "verbose", False): + if verbose: raise return 1 diff --git a/src/modelscope_hub/config.py b/src/modelscope_hub/config.py index 4d3056f..9fbbadc 100644 --- a/src/modelscope_hub/config.py +++ b/src/modelscope_hub/config.py @@ -38,6 +38,14 @@ ENV_TOKEN = "MODELSCOPE_API_TOKEN" ENV_HOME = "MODELSCOPE_HOME" +# Files that together constitute a persisted login. ``session`` is deliberately +# excluded: it is an anonymous SDK install identifier, not a credential. +_CREDENTIAL_FILE_NAMES: tuple[str, ...] = ( + COOKIES_FILE_NAME, + GIT_TOKEN_FILE_NAME, + USER_INFO_FILE_NAME, +) + def _expand(path: str | os.PathLike[str]) -> Path: return Path(path).expanduser().resolve() @@ -85,10 +93,7 @@ def __post_init__(self) -> None: self._endpoint_overridden = True else: self.endpoint = DEFAULT_ENDPOINT - # Ensure endpoint always has a scheme - if self.endpoint and not self.endpoint.startswith(("http://", "https://")): - self.endpoint = f"https://{self.endpoint}" - self.endpoint = (self.endpoint or DEFAULT_ENDPOINT).rstrip("/") + self.endpoint = self.normalize_endpoint(self.endpoint) # 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 @@ -101,6 +106,25 @@ def __post_init__(self) -> None: else: self.token = self.load_token() + @staticmethod + def normalize_endpoint(endpoint: str | None) -> str: + """Return *endpoint* with a scheme guaranteed and no trailing slash. + + Bare domains such as ``modelscope.ai`` are common input, especially from + the CLI. Without a scheme every request built from them fails deep in + the transport layer instead of surfacing a usable error, so the + normalisation lives here and is reused by every entry point that + accepts an endpoint. + + Scheme detection is case-insensitive because URI schemes are, so an + input like ``HTTPS://host`` is recognised instead of being prefixed a + second time. + """ + value = (endpoint or "").strip() or DEFAULT_ENDPOINT + if not value.lower().startswith(("http://", "https://")): + value = f"https://{value}" + return value.rstrip("/") + # ------------------------------------------------------------------ # Path helpers # ------------------------------------------------------------------ @@ -186,14 +210,19 @@ def load_token(self) -> str | None: return None def clear_token(self) -> None: - """Remove persisted credentials (deletes ``credentials/cookies``).""" + """Remove every persisted credential artefact. + + All login artefacts are dropped together. Removing only the session + cookie would leave the git token and the cached identity behind, a + half-logged-out state that later reads can still pick up. + """ self.token = None self._logged_out = True - path = self.credentials_dir / COOKIES_FILE_NAME - try: - path.unlink(missing_ok=True) - except OSError: - pass + for name in _CREDENTIAL_FILE_NAMES: + try: + (self.credentials_dir / name).unlink(missing_ok=True) + except OSError: + pass # ------------------------------------------------------------------ # Credentials persistence (compat with old modelscope SDK) diff --git a/src/modelscope_hub/errors.py b/src/modelscope_hub/errors.py index b4c9db3..52a685b 100644 --- a/src/modelscope_hub/errors.py +++ b/src/modelscope_hub/errors.py @@ -394,6 +394,40 @@ class NotSupportedError(HubError): } +# --------------------------------------------------------------------------- +# Server business-code -> exception mapping +# +# The HTTP status is not always faithful to the failure semantics: the legacy +# login endpoint answers 400 while meaning "authentication failed". Where the +# server publishes a business code, trust it over the status code. Register new +# codes in this table (and in the ModelScope error-code spec) rather than +# branching at the call site. +# --------------------------------------------------------------------------- +_BUSINESS_CODE_MAP: dict[int, type[APIError]] = { + # -> E3001, served with HTTP 400 by POST /api/v1/login on both sites + 10010103009: AuthenticationError, # AccessToken 无效或过期 + # -> E3026 + 10020101001: AlreadyExistsError, # 国内站 - 数据集已存在 + 10010101001: AlreadyExistsError, # 国内站 - 模型已存在 + 10010202004: AlreadyExistsError, # 国际站 - 名称已被使用 +} + + +def _business_code(body: Any) -> int | None: + """Return the numeric business code carried by a response body, if any.""" + if not isinstance(body, dict): + return None + raw = body.get("Code") + if raw is None: + raw = body.get("code") + if raw is None: + return None + try: + return int(raw) + except (TypeError, ValueError): + return None + + _CN_TO_EN: dict[str, str] = { "该名称已被注册使用,请重新命名": "Repository name already exists. Please choose a different name.", "用户未登录": "User not logged in.", @@ -479,21 +513,20 @@ def raise_for_status(response: Response) -> None: else: exc_cls = _STATUS_MAP.get(status, APIError) - # Detect "already exists" errors before falling back to InvalidParameter - if exc_cls is InvalidParameter and isinstance(body, dict): - code = body.get("Code") or body.get("code") + # A published business code describes a client-side failure more faithfully + # than the HTTP status, so it wins for 4xx. It deliberately does not apply to + # 5xx: a server outage must stay retryable even if the body happens to carry + # a known code, and reclassifying it would silently drop that retryability. + business_cls: type[APIError] | None = None + if status < 500: + code = _business_code(body) + business_cls = _BUSINESS_CODE_MAP.get(code) if code is not None else None + if business_cls is not None: + exc_cls = business_cls + elif exc_cls is InvalidParameter and isinstance(body, dict): + # Older servers signal "already exists" through the message only. msg_text = (body.get("Message") or body.get("message") or body.get("msg") or body.get("Msg") or "").lower() - is_exists = False - if code is not None: - try: - if int(code) in _ALREADY_EXISTS_CODES: - is_exists = True - except (TypeError, ValueError): - pass - if not is_exists: - if any(kw in msg_text for kw in _ALREADY_EXISTS_KEYWORDS): - is_exists = True - if is_exists: + if any(kw in msg_text for kw in _ALREADY_EXISTS_KEYWORDS): exc_cls = AlreadyExistsError kwargs: dict[str, Any] = dict( @@ -523,11 +556,12 @@ def raise_for_status(response: Response) -> None: # --------------------------------------------------------------------------- # Repo-exists detection (shared by cli/repo.py and compat/hub_api.py) # --------------------------------------------------------------------------- -_ALREADY_EXISTS_CODES: set[int] = { - 10020101001, # 国内站 - 数据集已存在 - 10010101001, # 国内站 - 模型已存在 - 10010202004, # 国际站 - 名称已被使用 -} +# Derived from the business-code table so the two never drift apart. Retained +# as a module-level name because :func:`is_repo_exists_error` still consults it +# when handling exceptions that pre-date the structured hierarchy. +_ALREADY_EXISTS_CODES: frozenset[int] = frozenset( + code for code, exc in _BUSINESS_CODE_MAP.items() if exc is AlreadyExistsError +) _ALREADY_EXISTS_KEYWORDS: frozenset[str] = frozenset( { diff --git a/tests/cli/test_main.py b/tests/cli/test_main.py index 4abce64..96e51a7 100644 --- a/tests/cli/test_main.py +++ b/tests/cli/test_main.py @@ -7,7 +7,7 @@ import pytest -from modelscope_hub.cli.main import run_cmd +from modelscope_hub.cli.main import _report_hub_error, run_cmd from modelscope_hub.errors import HubError, InvalidParameter, NetworkError, NotSupportedError from .conftest import run_cli @@ -85,6 +85,82 @@ def test_global_flags_before_subcommand(self, parser): assert args.verbose is True +# --------------------------------------------------------------------------- +# Error cause reporting (pure unit -- no API, so it runs in every mode) +# --------------------------------------------------------------------------- +class TestErrorCauseReporting: + """``--verbose`` must expose the failure that a wrapper exception hides. + + Without this, a wrapped SDK error showed only the outermost message and the + originating cause was unreachable from the CLI in any mode. + """ + + @staticmethod + def _wrapped_error() -> HubError: + """Build a three-level chain the way the SDK layers do.""" + try: + try: + try: + raise ValueError("socket closed") + except ValueError as root: + raise NetworkError("connection refused") from root + except NetworkError as mid: + raise HubError("outer failure") from mid + except HubError as exc: + return exc + + def test_cause_chain_hidden_without_verbose(self, capsys): + _report_hub_error(self._wrapped_error(), verbose=False) + captured = capsys.readouterr() + assert "outer failure" in captured.err + assert "Caused by" not in captured.out + + def test_cause_chain_shown_with_verbose(self, capsys): + _report_hub_error(self._wrapped_error(), verbose=True) + out = capsys.readouterr().out + # SDK errors render their own error code, so match on class + message. + assert "Caused by: NetworkError:" in out + assert "connection refused" in out + assert "Caused by: ValueError: socket closed" in out + + def test_suppressed_context_is_respected(self, capsys): + """``raise ... from None`` deliberately hides the context.""" + try: + try: + raise ValueError("hidden detail") + except ValueError: + raise HubError("clean failure") from None + except HubError as exc: + _report_hub_error(exc, verbose=True) + captured = capsys.readouterr() + assert "clean failure" in captured.err + assert "Caused by" not in captured.out + + def test_self_referential_chain_terminates(self, capsys): + """A cyclic chain must not stall the error path.""" + exc = HubError("looping failure") + exc.__cause__ = exc + + _report_hub_error(exc, verbose=True) + + assert "looping failure" in capsys.readouterr().err + + def test_chain_depth_is_bounded(self, capsys): + """Only the first *max_depth* causes are rendered.""" + exc = HubError("level-0") + current: BaseException = exc + for level in range(1, 8): + nested = ValueError(f"level-{level}") + current.__cause__ = nested + current = nested + + _report_hub_error(exc, verbose=True, max_depth=3) + + out = capsys.readouterr().out + assert out.count("Caused by") == 3 + assert "level-4" not in out + + # --------------------------------------------------------------------------- # Exception handling (unit tests with mocks — no API needed) # --------------------------------------------------------------------------- diff --git a/tests/test_credential_lifecycle.py b/tests/test_credential_lifecycle.py new file mode 100644 index 0000000..9c87cfe --- /dev/null +++ b/tests/test_credential_lifecycle.py @@ -0,0 +1,112 @@ +# Copyright (c) Alibaba, Inc. and its affiliates. +"""Regression guards for credential teardown. + +Covers two defects that left the SDK half-authenticated: + +* ``clear_token`` deleted only the session cookie, so the git token and the + cached identity outlived the login they belonged to. +* :attr:`HubApi.legacy` refused to propagate a *cleared* token to an already + constructed client, which then kept authenticating with the revoked + credential. +""" + +from __future__ import annotations + +import pytest + +from modelscope_hub.api import HubApi +from modelscope_hub.config import _CREDENTIAL_FILE_NAMES, HubConfig +from modelscope_hub.constants import ( + COOKIES_FILE_NAME, + GIT_TOKEN_FILE_NAME, + SESSION_FILE_NAME, + USER_INFO_FILE_NAME, +) + +TOKEN = "ms-token-under-test" +GIT_TOKEN = "git-token-value" + + +@pytest.fixture(autouse=True) +def isolated_home(tmp_path, monkeypatch): + """Redirect credential storage and drop ambient endpoint/token overrides.""" + for name in ("MODELSCOPE_ENDPOINT", "MODELSCOPE_API_TOKEN", "MODELSCOPE_DOMAIN"): + monkeypatch.delenv(name, raising=False) + monkeypatch.setenv("MODELSCOPE_HOME", str(tmp_path)) + return tmp_path + + +def fully_logged_in(home) -> HubConfig: + """Write every artefact that a successful login leaves behind.""" + config = HubConfig(config_dir=home) + config.save_token(TOKEN) + config.save_git_token(GIT_TOKEN) + config.save_user_info("alice", "alice@example.com") + config.get_session_id() # materialises the session file + for name in (COOKIES_FILE_NAME, GIT_TOKEN_FILE_NAME, USER_INFO_FILE_NAME, SESSION_FILE_NAME): + assert (config.credentials_dir / name).exists(), name + return config + + +# --------------------------------------------------------------------------- +# All-or-nothing teardown +# --------------------------------------------------------------------------- +def test_clear_token_removes_every_credential_artefact(isolated_home): + """A partial wipe used to leave a working git token behind.""" + config = fully_logged_in(isolated_home) + + config.clear_token() + + for name in _CREDENTIAL_FILE_NAMES: + assert not (config.credentials_dir / name).exists(), name + assert config.load_token() is None + assert config.load_git_token() is None + + +def test_clear_token_keeps_the_anonymous_session_id(isolated_home): + """The session id is an install identifier, not a credential.""" + config = fully_logged_in(isolated_home) + session_before = config.get_session_id() + + config.clear_token() + + assert (config.credentials_dir / SESSION_FILE_NAME).exists() + assert config.get_session_id() == session_before + + +def test_logout_clears_persisted_state(isolated_home): + """``HubApi.logout`` goes through the same all-or-nothing teardown.""" + fully_logged_in(isolated_home) + api = HubApi(config=HubConfig(config_dir=isolated_home)) + assert api._config.token == TOKEN + + api.logout() + + reloaded = HubConfig(config_dir=isolated_home) + assert reloaded.load_token() is None + assert reloaded.load_git_token() is None + + +# --------------------------------------------------------------------------- +# Cached client stays in step with the configured token +# --------------------------------------------------------------------------- +def test_cleared_token_propagates_to_the_cached_legacy_client(isolated_home): + """A cached client must not keep using a credential that was revoked.""" + fully_logged_in(isolated_home) + api = HubApi(config=HubConfig(config_dir=isolated_home)) + assert api.legacy.token == TOKEN # materialise the client + + api._config.clear_token() + + assert api.legacy.token is None + + +def test_rotated_token_propagates_to_the_cached_legacy_client(isolated_home): + """The pre-existing propagation path keeps working.""" + fully_logged_in(isolated_home) + api = HubApi(config=HubConfig(config_dir=isolated_home)) + assert api.legacy.token == TOKEN + + api._config.token = "ms-rotated" + + assert api.legacy.token == "ms-rotated" diff --git a/tests/test_login_failure_paths.py b/tests/test_login_failure_paths.py new file mode 100644 index 0000000..cfb880c --- /dev/null +++ b/tests/test_login_failure_paths.py @@ -0,0 +1,251 @@ +# Copyright (c) Alibaba, Inc. and its affiliates. +"""Regression guards for the ``HubApi.login`` failure paths. + +These tests stub the ``requests`` transport only, so the real +:class:`~modelscope_hub.config.HubConfig`, :class:`~modelscope_hub.api.HubApi`, +:class:`~modelscope_hub._legacy_api.LegacyClient` and error-translation layers +all take part. Mocking any higher would hide precisely the defects covered +here: a rejected login used to surface as a fabricated "token rejected" +message -- losing the server's own explanation and request id -- while +deleting the credential the caller already had on disk. + +The payloads below are the ones ``POST /api/v1/login`` actually returns. +""" + +from __future__ import annotations + +import json + +import pytest +import requests +import responses + +from modelscope_hub.api import HubApi +from modelscope_hub.config import HubConfig +from modelscope_hub.constants import DEFAULT_ENDPOINT, DEFAULT_INTL_ENDPOINT +from modelscope_hub.errors import ( + AlreadyExistsError, + AuthenticationError, + InvalidParameter, + NetworkError, + ServerError, + raise_for_status, +) + +CN_LOGIN = f"{DEFAULT_ENDPOINT}/api/v1/login" +AI_LOGIN = f"{DEFAULT_INTL_ENDPOINT}/api/v1/login" + +# Both sites answer an unknown token with this same business code, which is why +# the server cannot tell "invalid token" apart from "token issued elsewhere". +TOKEN_REJECTED_BODY = { + "Code": 10010103009, + "Message": "登录失败,AccessToken错误,请从用户中心获取AccessToken或刷新", + "RequestId": "8a039827-3f7c-4378-9b7c-3f8341b73649", + "Success": False, +} + +LOGIN_OK_BODY = { + "Code": 200, + "Data": {"AccessToken": "git-token", "Email": "alice@example.com", "Username": "alice"}, + "Message": "success", + "RequestId": "b008966a-942f-4c05-8f8a-696d1b6cc2e2", + "Success": True, +} + +PRIOR_TOKEN = "ms-previously-working" + + +@pytest.fixture(autouse=True) +def isolated_home(tmp_path, monkeypatch): + """Redirect credential storage and drop ambient endpoint/token overrides.""" + for name in ("MODELSCOPE_ENDPOINT", "MODELSCOPE_API_TOKEN", "MODELSCOPE_DOMAIN"): + monkeypatch.delenv(name, raising=False) + monkeypatch.setenv("MODELSCOPE_HOME", str(tmp_path)) + return tmp_path + + +def unpinned_api(home) -> HubApi: + """Build a HubApi with no explicit endpoint, so the peer probe may run.""" + return HubApi(config=HubConfig(config_dir=home)) + + +def seed_stored_credential(home) -> None: + """Persist a working credential the way a previous login would have.""" + config = HubConfig(config_dir=home) + config.save_token(PRIOR_TOKEN) + assert config.load_token() == PRIOR_TOKEN + + +def stored_token(home) -> str | None: + return HubConfig(config_dir=home).load_token() + + +# --------------------------------------------------------------------------- +# Faithful attribution +# --------------------------------------------------------------------------- +@responses.activate +def test_rejected_token_keeps_server_explanation(isolated_home): + """The server's message, business code and request id must survive.""" + responses.add(responses.POST, CN_LOGIN, json=TOKEN_REJECTED_BODY, status=400) + responses.add(responses.POST, AI_LOGIN, json=TOKEN_REJECTED_BODY, status=400) + + with pytest.raises(AuthenticationError) as excinfo: + unpinned_api(isolated_home).login("ms-bad-token") + + exc = excinfo.value + # Business code 10010103009 outranks the HTTP 400 the server chose. + assert exc.error_code == "E3001" + assert exc.status_code == 400 + assert exc.request_id == TOKEN_REJECTED_BODY["RequestId"] + assert TOKEN_REJECTED_BODY["Message"] in exc.message + + +@responses.activate +def test_network_failure_is_not_reported_as_bad_token(isolated_home): + """Transport errors keep their own identity instead of blaming the token.""" + seed_stored_credential(isolated_home) + responses.add(responses.POST, CN_LOGIN, body=requests.ConnectionError("connection reset")) + + with pytest.raises(NetworkError) as excinfo: + HubApi(config=HubConfig(config_dir=isolated_home)).login("ms-any-token") + + assert not isinstance(excinfo.value, AuthenticationError) + assert stored_token(isolated_home) == PRIOR_TOKEN + + +@responses.activate +def test_server_error_propagates_unchanged(isolated_home): + """A 5xx is a server outage, not a credential problem.""" + responses.add(responses.POST, CN_LOGIN, json={"Code": 500, "Message": "internal"}, status=500) + + with pytest.raises(ServerError) as excinfo: + unpinned_api(isolated_home).login("ms-any-token") + + assert excinfo.value.error_code == "E1002" + + +# --------------------------------------------------------------------------- +# Stored credentials are never collateral damage +# --------------------------------------------------------------------------- +@responses.activate +def test_failed_login_keeps_stored_credential(isolated_home): + """A mistyped token must not log the user out of a working session.""" + seed_stored_credential(isolated_home) + responses.add(responses.POST, CN_LOGIN, json=TOKEN_REJECTED_BODY, status=400) + responses.add(responses.POST, AI_LOGIN, json=TOKEN_REJECTED_BODY, status=400) + + api = HubApi(config=HubConfig(config_dir=isolated_home)) + with pytest.raises(AuthenticationError): + api.login("ms-bad-token") + + assert stored_token(isolated_home) == PRIOR_TOKEN + # The instance is also rewound, not left holding the rejected token. + assert api._config.token == PRIOR_TOKEN + + +# --------------------------------------------------------------------------- +# Site disambiguation +# --------------------------------------------------------------------------- +@responses.activate +def test_token_valid_on_peer_site_yields_endpoint_hint(isolated_home): + """A token issued by the other site gets an actionable hint, not a verdict.""" + responses.add(responses.POST, CN_LOGIN, json=TOKEN_REJECTED_BODY, status=400) + responses.add(responses.POST, AI_LOGIN, json=LOGIN_OK_BODY, status=200) + + with pytest.raises(AuthenticationError) as excinfo: + unpinned_api(isolated_home).login("ms-intl-token") + + message = excinfo.value.message + assert DEFAULT_INTL_ENDPOINT in message + assert "--endpoint" in message + + +@responses.activate +def test_pinned_endpoint_is_not_second_guessed(isolated_home): + """An explicit endpoint is respected: no peer probe and no hint.""" + responses.add(responses.POST, CN_LOGIN, json=TOKEN_REJECTED_BODY, status=400) + + api = HubApi(config=HubConfig(config_dir=isolated_home), endpoint=DEFAULT_ENDPOINT) + with pytest.raises(AuthenticationError) as excinfo: + api.login("ms-bad-token") + + assert "--endpoint" not in excinfo.value.message + assert all(AI_LOGIN not in call.request.url for call in responses.calls) + + +# --------------------------------------------------------------------------- +# Classification table +# --------------------------------------------------------------------------- +def json_response(status: int, body: dict) -> requests.Response: + """Build a minimal response the error layer can classify.""" + resp = requests.Response() + resp.status_code = status + resp._content = json.dumps(body).encode() + resp.headers["Content-Type"] = "application/json" + resp.url = CN_LOGIN + return resp + + +@pytest.mark.parametrize( + "body, expected", + [ + ({"Code": 10010103009, "Message": "token bad"}, AuthenticationError), + ({"Code": 10010101001, "Message": "model exists"}, AlreadyExistsError), + ({"Code": 99999999999, "Message": "something else"}, InvalidParameter), + ({"Message": "no code at all"}, InvalidParameter), + ], +) +def test_business_code_outranks_http_status(body, expected): + """A published business code classifies the failure; status is the fallback.""" + with pytest.raises(expected) as excinfo: + raise_for_status(json_response(400, body)) + # AlreadyExistsError subclasses InvalidParameter, so assert the exact type. + assert type(excinfo.value) is expected + + +@pytest.mark.parametrize("code", [10010103009, 10010101001]) +def test_business_code_does_not_override_a_server_outage(code): + """A 5xx stays a retryable ServerError even when the body carries a known code. + + Reclassifying it would flip ``retryable`` to False and silently turn a + transient outage into a permanent failure. + """ + with pytest.raises(ServerError) as excinfo: + raise_for_status(json_response(500, {"Code": code, "Message": "upstream failure"})) + + assert type(excinfo.value) is ServerError + assert excinfo.value.retryable is True + + +# --------------------------------------------------------------------------- +# Endpoint normalisation +# --------------------------------------------------------------------------- +@pytest.mark.parametrize( + "given, expected", + [ + ("modelscope.ai", "https://modelscope.ai"), + ("https://modelscope.cn/", "https://modelscope.cn"), + (" modelscope.cn ", "https://modelscope.cn"), + ("http://localhost:8080", "http://localhost:8080"), + # URI schemes are case-insensitive: recognise, do not prefix again. + ("HTTPS://modelscope.cn", "HTTPS://modelscope.cn"), + ("Http://localhost:8080", "Http://localhost:8080"), + ], +) +def test_bare_endpoint_gains_a_scheme(isolated_home, given, expected): + """Bare domains used to reach the transport layer and fail there.""" + api = HubApi(config=HubConfig(config_dir=isolated_home), endpoint=given) + assert api._config.endpoint == expected + + +@responses.activate +def test_bare_endpoint_reaches_the_expected_url(isolated_home): + """End-to-end proof that a scheme-less endpoint now resolves correctly.""" + url = "https://modelscope.ai/api/v1/login" + responses.add(responses.POST, url, json=TOKEN_REJECTED_BODY, status=400) + + api = HubApi(config=HubConfig(config_dir=isolated_home), endpoint="modelscope.ai") + with pytest.raises(AuthenticationError): + api.login("ms-bad-token") + + assert responses.calls[0].request.url == url From ce5c7b874f94a025c74b6f5bf1044f52c184ab4d Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E7=8F=AD=E6=89=AC?= Date: Sat, 1 Aug 2026 22:17:38 +0800 Subject: [PATCH 13/13] update news and bump version --- README.md | 5 +++++ src/modelscope_hub/version.py | 2 +- 2 files changed, 6 insertions(+), 1 deletion(-) diff --git a/README.md b/README.md index e33120c..f348a62 100644 --- a/README.md +++ b/README.md @@ -33,6 +33,11 @@ The official Python SDK & CLI for [ModelScope Hub](https://modelscope.cn) — do ## News +**v0.2.0** (2026-08-01) +- **Breaking**: `HubApi.login()` now raises `NetworkError` / `ServerError` / `RequestTimeoutError` for non-authentication failures instead of always raising `AuthenticationError` — widen `except AuthenticationError` to `HubError` if you catch *any* login failure +- **Fix**: network, timeout and 5xx errors during login are no longer misreported as an invalid token; a failed login no longer deletes stored credentials; bare and upper-case endpoints are accepted (e.g. `--endpoint modelscope.ai`) +- **Enhancement**: a rejected login probes the peer ModelScope site and suggests `--endpoint` when the token is valid there; `--verbose` prints the full error cause chain + **v0.1.9** (2026-07-31) - **Fix**: two stale unit tests that failed in downstream distro sandboxes ([#46](https://github.com/modelscope/modelscope_hub/issues/46), NixOS): align the `mcp deploy` call-shape assertion and the `parse_timestamp` timezone-normalization contract - **Feature**: `ms-hub mcp deploy` prints the operational URL after deployment; new `--auth-check` and repeatable `--env KEY=VALUE` options; `--transport-type` now validates against `sse`/`streamable_http` diff --git a/src/modelscope_hub/version.py b/src/modelscope_hub/version.py index 8804a1a..a1bc874 100644 --- a/src/modelscope_hub/version.py +++ b/src/modelscope_hub/version.py @@ -1,3 +1,3 @@ """Version information for modelscope_hub.""" -__version__ = "0.1.9+main" +__version__ = "0.2.0+main"