mirror of
https://github.com/willmiao/ComfyUI-Lora-Manager.git
synced 2026-09-20 18:51:26 -03:00
feat(download): support ModelScope repositories in the URL downloader
ModelScope became a linkable source, but downloading from it was impossible:
the URL picker only recognised huggingface.co, the file listing hit a
huggingface-only endpoint, the resolve URL was hardcoded, and the default
path template always wrote into a `huggingface/` directory.
Move the download knowledge into the providers so the handlers stay generic:
- `ModelSource` gains `list_files()`, `file_download_url()`,
`default_revision` and `default_subdir`. `HuggingFaceSource` keeps the Hub
tree API (`/api/models/{id}/tree/{rev}`, LFS-aware sizes, `main`).
`ModelScopeSource` uses `/api/v1/models/{id}/repo/files?Revision=master`
— which reports real byte sizes for LFS files, so no HEAD probe is needed,
and which only accepts `master` (an HF-imported repo still 404s on `main`)
— and downloads through `/models/{id}/resolve/{rev}/{path}`. That URL
redirects to a CDN target carrying a time-limited `auth_key`, so it is
rebuilt on every request and never cached, which is also what keeps
resumable Range requests working.
- `hf_handlers.py`/`HfHandler` become `model_source_handlers.py`/
`ModelSourceHandler` with `list_model_source_files` and
`download_model_source`. New routes `/api/lm/model-source-files` and
`/api/lm/download-model-source`; the old `/api/lm/hf-repo-files` and
`/api/lm/download-hf-model` paths stay as aliases, and a payload without
`platform` still means Hugging Face, so existing callers are unaffected.
- A downloaded sidecar now records `source_platform` + `source_url` (with the
`hf_url` alias only for Hugging Face) instead of always writing `hf_url`,
and `use_default_paths` files ModelScope downloads under
`modelscope/<owner>/<repo>`. The now-unused shared HF aiohttp session and
its shutdown hook are gone; providers open short-lived sessions.
- Frontend: `detectUrlType` returns the platform-neutral
`model-source-repo` / `model-source-file` plus an explicit `platform`, the
DownloadManager's `hf*` state and methods are renamed to `source*`, every
`source === 'huggingface'` check becomes `isExternalModelSource()`, and
batch groups are keyed by `platform:repo` so the same `owner/name` on two
sites renders as two groups. A bare `owner/name` still means Hugging Face.
- `is_valid_source_id()` centralises repo-id validation (exactly
`owner/name`, no traversal, no leading dot). This also fixes the old HF
download check that rejected any dot in the name, i.e. legitimate repos
such as `black-forest-labs/FLUX.1-dev`.
Verified against the live APIs: the example repo lists 8 weight files with
correct sizes, and a ranged GET of the built resolve URL returns 206 after
following the redirect to the CDN. Backend 2853 passed; frontend 1143 JS +
91 Vue passed. The nine locales carry the refreshed download copy in the
next commit.
This commit is contained in:
@@ -12,10 +12,14 @@ from .base import (
|
||||
GROUP_PREFIXES,
|
||||
HTTP_TIMEOUT,
|
||||
ModelSource,
|
||||
ModelSourceError,
|
||||
SourceRef,
|
||||
USER_AGENT,
|
||||
clean_source_url,
|
||||
fetch_json,
|
||||
fetch_text,
|
||||
filter_weight_files,
|
||||
is_valid_source_id,
|
||||
)
|
||||
from .huggingface import HuggingFaceSource
|
||||
from .modelscope import ModelScopeSource
|
||||
@@ -24,6 +28,8 @@ from .registry import (
|
||||
SOURCE_PLATFORM_FIELD,
|
||||
SOURCE_URL_FIELD,
|
||||
detect_source,
|
||||
downloadable_sources,
|
||||
get_download_source,
|
||||
get_source,
|
||||
get_source_platform,
|
||||
has_external_source,
|
||||
@@ -40,6 +46,7 @@ __all__ = [
|
||||
"HTTP_TIMEOUT",
|
||||
"LEGACY_HF_URL_FIELD",
|
||||
"ModelSource",
|
||||
"ModelSourceError",
|
||||
"HuggingFaceSource",
|
||||
"ModelScopeSource",
|
||||
"SOURCE_PLATFORM_FIELD",
|
||||
@@ -49,10 +56,15 @@ __all__ = [
|
||||
"USER_AGENT",
|
||||
"clean_source_url",
|
||||
"detect_source",
|
||||
"downloadable_sources",
|
||||
"fetch_json",
|
||||
"fetch_text",
|
||||
"filter_weight_files",
|
||||
"get_download_source",
|
||||
"get_source",
|
||||
"get_source_platform",
|
||||
"has_external_source",
|
||||
"is_valid_source_id",
|
||||
"list_sources",
|
||||
"normalize_metadata_source",
|
||||
"resolve_source_ref",
|
||||
|
||||
@@ -20,12 +20,15 @@ HTTP handlers never need site-specific branching.
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import os
|
||||
import re
|
||||
from dataclasses import dataclass
|
||||
from typing import Any, Optional
|
||||
from typing import Any, Iterable, Optional
|
||||
|
||||
import aiohttp
|
||||
|
||||
from ...utils.constants import MODEL_FILE_EXTENSIONS
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
#: Shared HTTP timeout for model-card fetches.
|
||||
@@ -58,6 +61,36 @@ class SourceRef:
|
||||
"""Canonical URL of the model page."""
|
||||
|
||||
|
||||
class ModelSourceError(Exception):
|
||||
"""Raised when a model source cannot satisfy a request.
|
||||
|
||||
Carries the HTTP status the API handler should answer with, so the
|
||||
handlers stay free of per-site error mapping.
|
||||
"""
|
||||
|
||||
def __init__(self, message: str, status: int = 502) -> None:
|
||||
super().__init__(message)
|
||||
self.status = status
|
||||
|
||||
|
||||
#: Repository ids are always exactly ``owner/name``. Components may contain
|
||||
#: dots (``black-forest-labs/FLUX.1-dev``) but must not be empty, ``.`` / ``..``,
|
||||
#: or start with a dot - the id is used as a path segment on disk.
|
||||
_SOURCE_ID_COMPONENT = re.compile(r"^[A-Za-z0-9_][A-Za-z0-9_.\-]*$")
|
||||
|
||||
|
||||
def is_valid_source_id(source_id: str) -> bool:
|
||||
"""Return ``True`` when *source_id* is a safe ``owner/name`` repository id."""
|
||||
|
||||
if not source_id or not isinstance(source_id, str) or source_id.count("/") != 1:
|
||||
return False
|
||||
owner, name = source_id.split("/", 1)
|
||||
return all(
|
||||
part and part not in (".", "..") and _SOURCE_ID_COMPONENT.match(part)
|
||||
for part in (owner, name)
|
||||
)
|
||||
|
||||
|
||||
async def fetch_text(url: str, *, timeout: int = HTTP_TIMEOUT) -> str:
|
||||
"""Fetch *url* and return its body as text, or ``""`` on any failure.
|
||||
|
||||
@@ -80,6 +113,34 @@ async def fetch_text(url: str, *, timeout: int = HTTP_TIMEOUT) -> str:
|
||||
return ""
|
||||
|
||||
|
||||
async def fetch_json(
|
||||
url: str, *, timeout: int = HTTP_TIMEOUT
|
||||
) -> tuple[int, Any]:
|
||||
"""Fetch *url* and return ``(status, parsed_body)``.
|
||||
|
||||
Unlike :func:`fetch_text` this reports the status, because callers such as
|
||||
the file-listing endpoints need to distinguish "repo not found" (404) from
|
||||
a transport failure. ``parsed_body`` is ``None`` when the response is not
|
||||
JSON or the request failed outright (status ``0``).
|
||||
"""
|
||||
|
||||
try:
|
||||
async with aiohttp.ClientSession(
|
||||
headers={"User-Agent": USER_AGENT},
|
||||
timeout=aiohttp.ClientTimeout(total=timeout),
|
||||
) as session:
|
||||
async with session.get(url) as resp:
|
||||
if resp.status != 200:
|
||||
return resp.status, None
|
||||
try:
|
||||
return resp.status, await resp.json(content_type=None)
|
||||
except Exception:
|
||||
return resp.status, None
|
||||
except Exception as exc: # pragma: no cover - network dependent
|
||||
logger.debug("Failed to fetch %s: %s", url, exc)
|
||||
return 0, None
|
||||
|
||||
|
||||
class ModelSource:
|
||||
"""Description and I/O for one external model hosting site."""
|
||||
|
||||
@@ -95,6 +156,12 @@ class ModelSource:
|
||||
#: Whether models can be downloaded directly from this site.
|
||||
supports_download: bool = False
|
||||
|
||||
#: Branch used when the caller does not pass an explicit revision.
|
||||
default_revision: str = ""
|
||||
|
||||
#: Sub-directory the "use default paths" template places downloads in.
|
||||
default_subdir: str = ""
|
||||
|
||||
#: Lenient pattern used to recognise URLs already stored in metadata.
|
||||
#: Captures the site-specific source id in group ``id``.
|
||||
url_pattern: re.Pattern[str] | None = None
|
||||
@@ -163,6 +230,45 @@ class ModelSource:
|
||||
|
||||
return ""
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# Download support
|
||||
# ------------------------------------------------------------------
|
||||
|
||||
async def list_files(
|
||||
self, source_id: str, revision: str = ""
|
||||
) -> list[dict[str, Any]]:
|
||||
"""List downloadable weight files in *source_id*.
|
||||
|
||||
Returns ``[{"filename": <repo-relative path>, "size": <bytes>}]``,
|
||||
largest first, filtered to :data:`MODEL_FILE_EXTENSIONS`. Sites
|
||||
without download support return an empty list.
|
||||
|
||||
Raises :class:`ModelSourceError` when the repository cannot be read,
|
||||
so the handler can surface "not found" separately from a transport
|
||||
failure.
|
||||
"""
|
||||
|
||||
return []
|
||||
|
||||
def file_download_url(
|
||||
self, source_id: str, filename: str, revision: str = ""
|
||||
) -> str:
|
||||
"""Return the direct (redirecting) download URL for one file."""
|
||||
|
||||
raise ModelSourceError(
|
||||
f"{self.label or self.platform} does not support downloads", status=400
|
||||
)
|
||||
|
||||
def resolve_revision(self, revision: str = "") -> str:
|
||||
"""Return *revision*, falling back to this site's default branch."""
|
||||
|
||||
return revision or self.default_revision
|
||||
|
||||
def page_url_for_file(self, source_id: str, filename: str) -> str:
|
||||
"""Return the human-facing page for *filename* inside *source_id*."""
|
||||
|
||||
return self.canonical_url(source_id)
|
||||
|
||||
def __repr__(self) -> str: # pragma: no cover - debugging aid
|
||||
return f"<ModelSource {self.platform}>"
|
||||
|
||||
@@ -175,12 +281,33 @@ def clean_source_url(url: Any) -> str:
|
||||
return url.strip()
|
||||
|
||||
|
||||
def filter_weight_files(entries: Iterable[tuple[str, int]]) -> list[dict[str, Any]]:
|
||||
"""Keep model-weight files from ``(path, size)`` pairs, largest first.
|
||||
|
||||
Every site lists a lot more than weights (READMEs, configs, tokenizers,
|
||||
…); the download picker only ever wants the files ComfyUI can load, which
|
||||
is exactly :data:`MODEL_FILE_EXTENSIONS`.
|
||||
"""
|
||||
|
||||
files = [
|
||||
{"filename": path, "size": int(size or 0)}
|
||||
for path, size in entries
|
||||
if path and os.path.splitext(path)[1].lower() in MODEL_FILE_EXTENSIONS
|
||||
]
|
||||
files.sort(key=lambda entry: entry["size"], reverse=True)
|
||||
return files
|
||||
|
||||
|
||||
__all__ = [
|
||||
"GROUP_PREFIXES",
|
||||
"HTTP_TIMEOUT",
|
||||
"ModelSource",
|
||||
"ModelSourceError",
|
||||
"SourceRef",
|
||||
"USER_AGENT",
|
||||
"clean_source_url",
|
||||
"fetch_json",
|
||||
"fetch_text",
|
||||
"filter_weight_files",
|
||||
"is_valid_source_id",
|
||||
]
|
||||
|
||||
@@ -2,9 +2,18 @@
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import re
|
||||
|
||||
from .base import ModelSource, fetch_text
|
||||
from .base import (
|
||||
ModelSource,
|
||||
ModelSourceError,
|
||||
fetch_json,
|
||||
fetch_text,
|
||||
filter_weight_files,
|
||||
)
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
#: Lenient — used to normalise URLs already stored in metadata; tolerates
|
||||
#: sub-paths such as ``/resolve/main/model.safetensors``.
|
||||
@@ -25,6 +34,8 @@ class HuggingFaceSource(ModelSource):
|
||||
label = "Hugging Face"
|
||||
supports_enrichment = True
|
||||
supports_download = True
|
||||
default_revision = "main"
|
||||
default_subdir = "huggingface"
|
||||
url_pattern = _URL_PATTERN
|
||||
strict_url_pattern = _STRICT_URL_PATTERN
|
||||
|
||||
@@ -32,7 +43,7 @@ class HuggingFaceSource(ModelSource):
|
||||
return f"https://huggingface.co/{source_id}"
|
||||
|
||||
def asset_base_url(self, source_id: str, revision: str = "") -> str:
|
||||
return f"https://huggingface.co/{source_id}/resolve/{revision or 'main'}"
|
||||
return f"https://huggingface.co/{source_id}/resolve/{self.resolve_revision(revision)}"
|
||||
|
||||
async def fetch_model_card(self, source_id: str) -> str:
|
||||
"""Fetch ``README.md`` from Hugging Face (tries ``main``, then ``master``)."""
|
||||
@@ -45,5 +56,51 @@ class HuggingFaceSource(ModelSource):
|
||||
return text
|
||||
return ""
|
||||
|
||||
async def list_files(
|
||||
self, source_id: str, revision: str = ""
|
||||
) -> list[dict]:
|
||||
"""List weight files via the Hub tree API.
|
||||
|
||||
The tree endpoint (rather than the model-info endpoint) is used
|
||||
because it reports accurate sizes for LFS-tracked files.
|
||||
"""
|
||||
|
||||
revision = self.resolve_revision(revision)
|
||||
status, payload = await fetch_json(
|
||||
f"https://huggingface.co/api/models/{source_id}/tree/{revision}"
|
||||
)
|
||||
|
||||
if status == 404:
|
||||
raise ModelSourceError(f"Repository '{source_id}' not found", status=404)
|
||||
if status != 200 or not isinstance(payload, list):
|
||||
raise ModelSourceError(
|
||||
f"Hugging Face API error while listing '{source_id}' (HTTP {status})"
|
||||
)
|
||||
|
||||
entries = []
|
||||
for entry in payload:
|
||||
if not isinstance(entry, dict):
|
||||
continue
|
||||
path = entry.get("path", "")
|
||||
size = entry.get("size", 0) or 0
|
||||
if not size and isinstance(entry.get("lfs"), dict):
|
||||
size = entry["lfs"].get("size", 0) or 0
|
||||
entries.append((path, size))
|
||||
|
||||
return filter_weight_files(entries)
|
||||
|
||||
def file_download_url(
|
||||
self, source_id: str, filename: str, revision: str = ""
|
||||
) -> str:
|
||||
return (
|
||||
f"https://huggingface.co/{source_id}/resolve/"
|
||||
f"{self.resolve_revision(revision)}/{filename}"
|
||||
)
|
||||
|
||||
def page_url_for_file(self, source_id: str, filename: str) -> str:
|
||||
return (
|
||||
f"https://huggingface.co/{source_id}/blob/{self.default_revision}/{filename}"
|
||||
)
|
||||
|
||||
|
||||
__all__ = ["HuggingFaceSource"]
|
||||
|
||||
@@ -2,20 +2,38 @@
|
||||
|
||||
ModelScope exposes the same "model card as README.md" convention as
|
||||
Hugging Face, including a YAML frontmatter block that often carries
|
||||
``base_model:`` and ``trigger_words:``. Two public endpoints are used,
|
||||
neither of which requires an API key for public models:
|
||||
``base_model:`` and ``trigger_words:``. Three public endpoints are used,
|
||||
none of which requires an API key for public models:
|
||||
|
||||
* ``/models/{owner}/{name}/resolve/{revision}/README.md`` — raw model card
|
||||
* ``/api/v1/models/{owner}/{name}/repo?Revision=..&FilePath=README.md`` —
|
||||
the same content through the API, used as a fallback when the resolve
|
||||
URL is unavailable.
|
||||
* ``/api/v1/models/{owner}/{name}/repo/files?Revision=..`` — the file
|
||||
listing backing the download picker. It reports real sizes for LFS
|
||||
files (not the pointer size), so no extra HEAD request is needed.
|
||||
|
||||
Downloads go through ``/models/{owner}/{name}/resolve/{revision}/{path}``,
|
||||
which redirects to a CDN URL carrying a time-limited ``auth_key``.
|
||||
Requesting the resolve URL fresh on every attempt (which the shared
|
||||
downloader does, including for resumable Range requests) keeps that key
|
||||
valid; the CDN URL must never be cached.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import re
|
||||
|
||||
from .base import ModelSource, fetch_text
|
||||
from .base import (
|
||||
ModelSource,
|
||||
ModelSourceError,
|
||||
fetch_json,
|
||||
fetch_text,
|
||||
filter_weight_files,
|
||||
)
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
_URL_PATTERN = re.compile(
|
||||
r"https?://(?:www\.)?modelscope\.(?:cn|com)/models/(?P<id>[^/?#\s]+/[^/?#\s]+)"
|
||||
@@ -41,7 +59,9 @@ class ModelScopeSource(ModelSource):
|
||||
platform = "modelscope"
|
||||
label = "ModelScope"
|
||||
supports_enrichment = True
|
||||
supports_download = False
|
||||
supports_download = True
|
||||
default_revision = "master"
|
||||
default_subdir = "modelscope"
|
||||
url_pattern = _URL_PATTERN
|
||||
strict_url_pattern = _STRICT_URL_PATTERN
|
||||
|
||||
@@ -49,7 +69,10 @@ class ModelScopeSource(ModelSource):
|
||||
return f"https://modelscope.cn/models/{source_id}"
|
||||
|
||||
def asset_base_url(self, source_id: str, revision: str = "") -> str:
|
||||
return f"https://modelscope.cn/models/{source_id}/resolve/{revision or 'master'}"
|
||||
return (
|
||||
f"https://modelscope.cn/models/{source_id}/resolve/"
|
||||
f"{self.resolve_revision(revision)}"
|
||||
)
|
||||
|
||||
async def fetch_model_card(self, source_id: str) -> str:
|
||||
"""Fetch the model card, preferring the raw resolve URL."""
|
||||
@@ -72,5 +95,50 @@ class ModelScopeSource(ModelSource):
|
||||
return text
|
||||
return ""
|
||||
|
||||
async def list_files(
|
||||
self, source_id: str, revision: str = ""
|
||||
) -> list[dict]:
|
||||
"""List weight files via the repo files API.
|
||||
|
||||
``master`` is the only branch name the API accepts — even repos
|
||||
imported from Hugging Face are addressed as ``master`` (``main``
|
||||
returns 404) — so no fallback probing is done here.
|
||||
"""
|
||||
|
||||
revision = self.resolve_revision(revision)
|
||||
status, payload = await fetch_json(
|
||||
"https://modelscope.cn/api/v1/models/"
|
||||
f"{source_id}/repo/files?Revision={revision}"
|
||||
)
|
||||
|
||||
if status == 404:
|
||||
raise ModelSourceError(f"Repository '{source_id}' not found", status=404)
|
||||
if status != 200 or not isinstance(payload, dict):
|
||||
raise ModelSourceError(
|
||||
f"ModelScope API error while listing '{source_id}' (HTTP {status})"
|
||||
)
|
||||
|
||||
entries = []
|
||||
for entry in (payload.get("Data") or {}).get("Files") or []:
|
||||
if not isinstance(entry, dict) or entry.get("Type") != "blob":
|
||||
continue
|
||||
entries.append((entry.get("Path", ""), entry.get("Size", 0) or 0))
|
||||
|
||||
return filter_weight_files(entries)
|
||||
|
||||
def file_download_url(
|
||||
self, source_id: str, filename: str, revision: str = ""
|
||||
) -> str:
|
||||
return (
|
||||
f"https://modelscope.cn/models/{source_id}/resolve/"
|
||||
f"{self.resolve_revision(revision)}/{filename}"
|
||||
)
|
||||
|
||||
def page_url_for_file(self, source_id: str, filename: str) -> str:
|
||||
return (
|
||||
f"https://modelscope.cn/models/{source_id}/file/view/"
|
||||
f"{self.default_revision}/{filename}"
|
||||
)
|
||||
|
||||
|
||||
__all__ = ["ModelScopeSource"]
|
||||
|
||||
@@ -56,6 +56,21 @@ def source_label(platform: Optional[str], default: str = "") -> str:
|
||||
return source.label if source else default
|
||||
|
||||
|
||||
def downloadable_sources() -> list[ModelSource]:
|
||||
"""Return the sources whose repositories can be downloaded directly."""
|
||||
|
||||
return [source for source in _SOURCES if source.supports_download]
|
||||
|
||||
|
||||
def get_download_source(platform: Optional[str]) -> Optional[ModelSource]:
|
||||
"""Return the source for *platform*, but only when it supports downloads."""
|
||||
|
||||
source = get_source(platform)
|
||||
if source is None or not source.supports_download:
|
||||
return None
|
||||
return source
|
||||
|
||||
|
||||
def detect_source(url: Optional[str], *, strict: bool = False) -> Optional[SourceRef]:
|
||||
"""Return the :class:`SourceRef` for *url*, or ``None`` if unsupported."""
|
||||
|
||||
@@ -197,6 +212,8 @@ __all__ = [
|
||||
"SOURCE_PLATFORM_FIELD",
|
||||
"SOURCE_URL_FIELD",
|
||||
"detect_source",
|
||||
"downloadable_sources",
|
||||
"get_download_source",
|
||||
"get_source",
|
||||
"get_source_platform",
|
||||
"has_external_source",
|
||||
|
||||
Reference in New Issue
Block a user