mirror of
https://github.com/willmiao/ComfyUI-Lora-Manager.git
synced 2026-10-04 00:55:32 -03:00
feat(metadata): add OpenModelDB metadata provider and model source for upscalers
Add OpenModelDB (openmodeldb.info) as a metadata and download source for
the existing upscaler model type.
Metadata:
- New OpenModelDBClient: fetches the site's bulk JSON dumps, caches them
on disk (24h TTL + ETag revalidation), and builds a local sha256 index
- New OpenModelDBModelMetadataProvider adapts catalogue entries to the
CivitAI-shaped version dict contract; registered in the fallback chain
behind the enable_openmodeldb_api setting (default on), gated to the
upscaler sub-type so other model types never trigger the dump download
- Persisted provenance uses metadata_source "openmodeldb" plus a nested
openmodeldb block (page URL, architecture, scale, license)
Images: paired-image LR/SR URLs are ephemeral imgdiff.net sessions, so
displayable images come from the site-hosted auto-generated thumbnails
(model-level cover leads images[], per-image thumbs for the rest); the
original comparison URL is kept in meta.comparisonUrl.
Downloads:
- New OpenModelDBSource (flat model ids, omdb: group prefix) with
resource filename derivation that recovers names hidden mid-path
(mediafire) or synthesizes {id}.{type} for folder links
- HTML-gateway mirrors (mediafire/mega/drive) are rejected with a clear
manual-download hint instead of silently saving an HTML page as .pth
- ModelSource base gains is_valid_source_id / default_subdir_parts /
resolve_download_url hooks so flat-id sources need no platform branches
UI: "View on OpenModelDB" link in the model modal (downloaded and
hash-enriched models), settings toggle next to the CivArchive one.
This commit is contained in:
@@ -10,6 +10,7 @@ from .model_metadata_provider import (
|
||||
SQLiteModelMetadataProvider,
|
||||
CivitaiModelMetadataProvider,
|
||||
CivArchiveModelMetadataProvider,
|
||||
OpenModelDBModelMetadataProvider,
|
||||
FallbackMetadataProvider,
|
||||
RateLimitRetryingProvider,
|
||||
)
|
||||
@@ -22,12 +23,18 @@ logger = logging.getLogger(__name__)
|
||||
_PROVIDER_DISPLAY_NAMES = {
|
||||
"civitai_api": "CivitAI",
|
||||
"civarchive_api": "CivArchive",
|
||||
"openmodeldb_api": "OpenModelDB",
|
||||
"sqlite": "Archive DB",
|
||||
}
|
||||
|
||||
# Preset fallback chains. civitai_api is always first (richest metadata).
|
||||
# openmodeldb_api sits right after it: its lookups are local index hits over a
|
||||
# cached bulk dump (no rate-limit budget spent), and it covers upscalers that
|
||||
# CivArchive only has when they once existed on CivitAI. Providers that are not
|
||||
# registered (disabled/unavailable) are skipped, so presets degrade gracefully.
|
||||
_PRESET_PROVIDER_ORDERS = {
|
||||
"civitai_archive_sqlite": ["civitai_api", "civarchive_api", "sqlite"],
|
||||
"civitai_sqlite_archive": ["civitai_api", "sqlite", "civarchive_api"],
|
||||
"civitai_archive_sqlite": ["civitai_api", "openmodeldb_api", "civarchive_api", "sqlite"],
|
||||
"civitai_sqlite_archive": ["civitai_api", "openmodeldb_api", "sqlite", "civarchive_api"],
|
||||
}
|
||||
|
||||
async def initialize_metadata_providers():
|
||||
@@ -42,6 +49,7 @@ async def initialize_metadata_providers():
|
||||
settings_manager = get_settings_manager()
|
||||
enable_archive_db = settings_manager.get('enable_metadata_archive_db', False)
|
||||
enable_civarchive_api = settings_manager.get('enable_civarchive_api', True)
|
||||
enable_openmodeldb_api = settings_manager.get('enable_openmodeldb_api', True)
|
||||
provider_order = settings_manager.get('metadata_provider_order', 'civitai_archive_sqlite')
|
||||
|
||||
providers = []
|
||||
@@ -92,6 +100,22 @@ async def initialize_metadata_providers():
|
||||
else:
|
||||
logger.debug("CivArchive metadata provider disabled by setting 'enable_civarchive_api'")
|
||||
|
||||
# Register the OpenModelDB provider when enabled. It only covers upscaler
|
||||
# models (hash-matched against its catalogue dump), so it complements
|
||||
# rather than replaces the CivitAI-family providers; disabling it avoids
|
||||
# the one-time bulk dump download entirely.
|
||||
if enable_openmodeldb_api:
|
||||
try:
|
||||
openmodeldb_client = await ServiceRegistry.get_openmodeldb_client()
|
||||
openmodeldb_provider = OpenModelDBModelMetadataProvider(openmodeldb_client)
|
||||
provider_manager.register_provider('openmodeldb_api', openmodeldb_provider)
|
||||
providers.append(('openmodeldb_api', openmodeldb_provider))
|
||||
logger.debug("OpenModelDB metadata provider registered (also included in fallback)")
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to initialize OpenModelDB metadata provider: {e}")
|
||||
else:
|
||||
logger.debug("OpenModelDB metadata provider disabled by setting 'enable_openmodeldb_api'")
|
||||
|
||||
# Preset fallback orderings (see module-level _PRESET_PROVIDER_ORDERS).
|
||||
# civitai_api is always first (better metadata); the remaining providers
|
||||
# are arranged by the configured preset. Providers that are not
|
||||
@@ -135,6 +159,7 @@ async def update_metadata_providers():
|
||||
settings_manager = get_settings_manager()
|
||||
enable_archive_db = settings_manager.get('enable_metadata_archive_db', False)
|
||||
enable_civarchive_api = settings_manager.get('enable_civarchive_api', True)
|
||||
enable_openmodeldb_api = settings_manager.get('enable_openmodeldb_api', True)
|
||||
provider_order = settings_manager.get('metadata_provider_order', 'civitai_archive_sqlite')
|
||||
|
||||
# Reinitialize all providers with new settings
|
||||
@@ -153,9 +178,10 @@ async def update_metadata_providers():
|
||||
)
|
||||
|
||||
logger.info(
|
||||
"Updated metadata providers: archive_db=%s, civarchive_api=%s, chain=%s",
|
||||
"Updated metadata providers: archive_db=%s, civarchive_api=%s, openmodeldb_api=%s, chain=%s",
|
||||
enable_archive_db,
|
||||
enable_civarchive_api,
|
||||
enable_openmodeldb_api,
|
||||
chain,
|
||||
)
|
||||
return provider_manager
|
||||
|
||||
@@ -15,11 +15,48 @@ from ..utils.models import autov3_from_civitai_files
|
||||
from ..utils.sidecar_paths import get_metadata_path
|
||||
from .connectivity_guard import OFFLINE_FRIENDLY_MESSAGE, is_expected_offline_error
|
||||
from .errors import RateLimitError
|
||||
from .model_sources import has_external_source
|
||||
from .model_metadata_provider import _LOCAL_PROVIDER_LABELS
|
||||
from .model_sources import get_source_platform, has_external_source
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
# Providers restricted to specific model sub_types, keyed by their
|
||||
# registration label. Providers not listed apply to every model type.
|
||||
# OpenModelDB only indexes upscalers, so it is never consulted for other model
|
||||
# types — that keeps its one-time bulk-catalogue download from being paid by
|
||||
# users who manage no upscalers at all.
|
||||
_PROVIDER_SUB_TYPE_RESTRICTIONS: Dict[str, frozenset] = {
|
||||
"openmodeldb_api": frozenset({"upscaler"}),
|
||||
}
|
||||
|
||||
#: External-source platforms that have their own hash-lookup metadata
|
||||
#: provider. A model downloaded from one of these is refreshed against the
|
||||
#: source's own catalogue first (upgrading the download-time card to the full
|
||||
#: payload) before CivitAI is consulted at all.
|
||||
_EXTERNAL_SOURCE_METADATA_PROVIDERS: Dict[str, str] = {
|
||||
"openmodeldb": "openmodeldb_api",
|
||||
}
|
||||
|
||||
|
||||
def _restricted_providers_for_sub_type(sub_type: Optional[str]) -> list:
|
||||
"""Return the restricted provider labels that apply to ``sub_type``."""
|
||||
return [
|
||||
name
|
||||
for name, allowed in _PROVIDER_SUB_TYPE_RESTRICTIONS.items()
|
||||
if sub_type in allowed
|
||||
]
|
||||
|
||||
|
||||
def _inapplicable_providers_for_sub_type(sub_type: Optional[str]) -> set:
|
||||
"""Return the restricted provider labels that do NOT apply to ``sub_type``."""
|
||||
return {
|
||||
name
|
||||
for name, allowed in _PROVIDER_SUB_TYPE_RESTRICTIONS.items()
|
||||
if sub_type not in allowed
|
||||
}
|
||||
|
||||
|
||||
def _merge_ordered_unique(existing: Iterable[str], new: Iterable[str]) -> list[str]:
|
||||
"""Concatenate two word lists, dropping duplicates without reordering.
|
||||
|
||||
@@ -226,6 +263,18 @@ class MetadataSyncService:
|
||||
sqlite_attempted = False
|
||||
|
||||
if model_data.get("civitai_deleted") is True:
|
||||
# Sub_type-restricted providers (e.g. OpenModelDB for
|
||||
# upscalers) stay reachable for deleted models: their
|
||||
# catalogues grow independently of CivitAI, so a model deleted
|
||||
# from CivitAI may still gain metadata there later.
|
||||
for restricted_name in _restricted_providers_for_sub_type(
|
||||
model_data.get("sub_type")
|
||||
):
|
||||
try:
|
||||
provider_attempts.append((restricted_name, await self._get_provider(restricted_name)))
|
||||
except Exception as exc: # pragma: no cover - provider resolution fault
|
||||
logger.debug("Unable to resolve %s provider: %s", restricted_name, exc)
|
||||
|
||||
if previous_source in (None, "civarchive"):
|
||||
try:
|
||||
provider_attempts.append(("civarchive_api", await self._get_provider("civarchive_api")))
|
||||
@@ -250,19 +299,44 @@ class MetadataSyncService:
|
||||
is_hf_source = has_external_source(model_data)
|
||||
if is_hf_source:
|
||||
# External-source model (Hugging Face / ModelScope /
|
||||
# TensorArt): only check CivitAI API directly.
|
||||
# CivArchive is almost guaranteed to have no record, and
|
||||
# hitting it wastes rate-limit budget.
|
||||
# TensorArt / OpenModelDB): a source with its own
|
||||
# hash-lookup provider (OpenModelDB) is consulted first,
|
||||
# then CivitAI API directly. CivArchive is almost
|
||||
# guaranteed to have no record, and hitting it wastes
|
||||
# rate-limit budget.
|
||||
# Use a distinct provider name ("civitai_api" not None) so
|
||||
# downstream code does NOT interpret a "Model not found"
|
||||
# response as civitai_api_not_found — which would mark the
|
||||
# model civitai_deleted=True when it was never on CivitAI.
|
||||
try:
|
||||
provider_attempts.append(("civitai_api", await self._get_provider("civitai_api")))
|
||||
except Exception as exc: # pragma: no cover - provider resolution fault
|
||||
logger.debug("Unable to resolve civitai_api provider: %s", exc)
|
||||
source_provider = _EXTERNAL_SOURCE_METADATA_PROVIDERS.get(
|
||||
get_source_platform(model_data)
|
||||
)
|
||||
provider_names = (
|
||||
[source_provider, "civitai_api"]
|
||||
if source_provider
|
||||
else ["civitai_api"]
|
||||
)
|
||||
for provider_name in provider_names:
|
||||
try:
|
||||
provider_attempts.append(
|
||||
(provider_name, await self._get_provider(provider_name))
|
||||
)
|
||||
except Exception as exc: # pragma: no cover - provider resolution fault
|
||||
logger.debug(
|
||||
"Unable to resolve %s provider: %s", provider_name, exc
|
||||
)
|
||||
if not provider_attempts:
|
||||
provider_attempts.append((None, await self._get_default_provider()))
|
||||
default_provider = await self._get_default_provider()
|
||||
# Drop sub_type-restricted providers that cannot apply to
|
||||
# this model (e.g. OpenModelDB only indexes upscalers), so
|
||||
# their cold-start cost is never paid pointlessly.
|
||||
inapplicable = _inapplicable_providers_for_sub_type(
|
||||
model_data.get("sub_type")
|
||||
)
|
||||
excluding = getattr(default_provider, "excluding", None)
|
||||
if inapplicable and callable(excluding):
|
||||
default_provider = excluding(inapplicable)
|
||||
provider_attempts.append((None, default_provider))
|
||||
|
||||
civitai_metadata: Optional[Dict[str, Any]] = None
|
||||
metadata_provider: Optional[MetadataProviderProtocol] = None
|
||||
@@ -273,10 +347,11 @@ class MetadataSyncService:
|
||||
|
||||
skip_network_providers = False
|
||||
for provider_name, provider in provider_attempts:
|
||||
if skip_network_providers and provider_name != "sqlite":
|
||||
if skip_network_providers and provider_name not in _LOCAL_PROVIDER_LABELS:
|
||||
# A network provider was already rate-limited; failing
|
||||
# over to another network provider just spreads the flood
|
||||
# (#1085). The local sqlite archive stays as last resort.
|
||||
# (#1085). Local lookups (sqlite archive, the cached
|
||||
# OpenModelDB index) stay available as a last resort.
|
||||
continue
|
||||
try:
|
||||
civitai_metadata_candidate, error = await provider.get_model_by_hash(sha256)
|
||||
@@ -386,6 +461,7 @@ class MetadataSyncService:
|
||||
readable_source = {
|
||||
"civitai_api": "CivitAI API",
|
||||
"civarchive": "CivArchive API",
|
||||
"openmodeldb": "OpenModelDB",
|
||||
"archive_db": "Archive Database",
|
||||
}.get(source, source)
|
||||
|
||||
|
||||
@@ -112,7 +112,10 @@ class _RateLimitRetryHelper:
|
||||
|
||||
# Labels of providers that are free to consult even while a network provider
|
||||
# is rate-limited (local lookups, no vendor cost).
|
||||
_LOCAL_PROVIDER_LABELS = frozenset({"sqlite"})
|
||||
# "openmodeldb_api" qualifies because its lookups hit a local index built from
|
||||
# a cached bulk dump; the underlying site is a static host (GitHub Pages), so
|
||||
# even a cold cache refresh is a single cheap GET against a different vendor.
|
||||
_LOCAL_PROVIDER_LABELS = frozenset({"sqlite", "openmodeldb_api"})
|
||||
|
||||
|
||||
class ModelMetadataProvider(ABC):
|
||||
@@ -480,6 +483,36 @@ class SQLiteModelMetadataProvider(ModelMetadataProvider):
|
||||
except json.JSONDecodeError:
|
||||
return None
|
||||
|
||||
class OpenModelDBModelMetadataProvider(ModelMetadataProvider):
|
||||
"""Provider that serves upscaler metadata from the OpenModelDB catalogue.
|
||||
|
||||
Only hash lookups are supported: OpenModelDB has no per-model or version
|
||||
API, so the remaining provider surface intentionally returns None and lets
|
||||
the fallback chain continue to the next provider.
|
||||
"""
|
||||
|
||||
def __init__(self, openmodeldb_client):
|
||||
self.client = openmodeldb_client
|
||||
|
||||
async def get_model_by_hash(self, model_hash: str) -> Tuple[Optional[Dict[str, Any]], Optional[str]]:
|
||||
return await self.client.get_model_by_hash(model_hash)
|
||||
|
||||
async def get_model_versions(self, model_id: str) -> Optional[Dict[str, Any]]:
|
||||
"""Not supported: OpenModelDB models have no version history API."""
|
||||
return None
|
||||
|
||||
async def get_model_version(self, model_id: Optional[int] = None, version_id: Optional[int] = None) -> Optional[Dict[str, Any]]:
|
||||
"""Not supported: OpenModelDB models have no version history API."""
|
||||
return None
|
||||
|
||||
async def get_model_version_info(self, version_id: str) -> Tuple[Optional[Dict[str, Any]], Optional[str]]:
|
||||
"""Not supported: OpenModelDB models have no version history API."""
|
||||
return None, "Model not found"
|
||||
|
||||
async def get_user_models(self, username: str, cursor: Optional[str] = None) -> Optional[Dict[str, Any]]:
|
||||
"""Not supported by the OpenModelDB provider."""
|
||||
return None
|
||||
|
||||
class FallbackMetadataProvider(ModelMetadataProvider):
|
||||
"""Try providers in order, return first successful result.
|
||||
|
||||
@@ -750,6 +783,32 @@ class FallbackMetadataProvider(ModelMetadataProvider):
|
||||
def _iter_providers(self):
|
||||
return zip(self.providers, self._provider_labels)
|
||||
|
||||
def excluding(self, labels: "frozenset[str] | set[str]") -> "FallbackMetadataProvider":
|
||||
"""Return a copy of this chain without the providers named in *labels*.
|
||||
|
||||
Used by the metadata sync service to skip providers that cannot apply
|
||||
to a given model (e.g. OpenModelDB only indexes upscalers), so their
|
||||
cold-start cost (a bulk dump download) is never paid pointlessly.
|
||||
"""
|
||||
kept = [
|
||||
(label, provider)
|
||||
for provider, label in self._iter_providers()
|
||||
if label not in labels
|
||||
]
|
||||
if len(kept) == len(self.providers):
|
||||
return self
|
||||
if not kept:
|
||||
# Never produce an empty chain; the caller still needs a provider
|
||||
# that can at least report "Model not found".
|
||||
return self
|
||||
return FallbackMetadataProvider(
|
||||
kept,
|
||||
rate_limit_retry_limit=self._rate_limit_retry_limit,
|
||||
rate_limit_base_delay=self._rate_limit_base_delay,
|
||||
rate_limit_max_delay=self._rate_limit_max_delay,
|
||||
rate_limit_jitter_ratio=self._rate_limit_jitter_ratio,
|
||||
)
|
||||
|
||||
async def _call_with_rate_limit(self, label: str, func, *args, **kwargs):
|
||||
return await self._rate_limit_helper.run(label, func, *args, **kwargs)
|
||||
|
||||
|
||||
@@ -1,4 +1,4 @@
|
||||
"""External model-source providers (Hugging Face, ModelScope, TensorArt).
|
||||
"""External model-source providers (Hugging Face, ModelScope, TensorArt, OpenModelDB).
|
||||
|
||||
This package is the single abstraction over "a site that hosts models and
|
||||
a model card". See :mod:`py.services.model_sources.base` for the provider
|
||||
@@ -30,6 +30,7 @@ from .hydration import (
|
||||
resolve_site_base_model,
|
||||
)
|
||||
from .modelscope import ModelScopeIntlSource, ModelScopeSource
|
||||
from .openmodeldb import OpenModelDBSource
|
||||
from .registry import (
|
||||
LEGACY_HF_URL_FIELD,
|
||||
SOURCE_PLATFORM_FIELD,
|
||||
@@ -59,6 +60,7 @@ __all__ = [
|
||||
"HuggingFaceSource",
|
||||
"ModelScopeIntlSource",
|
||||
"ModelScopeSource",
|
||||
"OpenModelDBSource",
|
||||
"SOURCE_PLATFORM_FIELD",
|
||||
"SOURCE_URL_FIELD",
|
||||
"SourceRef",
|
||||
|
||||
@@ -47,6 +47,7 @@ GROUP_PREFIXES: dict[str, str] = {
|
||||
"modelscope": "ms",
|
||||
"modelscope-ai": "msai",
|
||||
"tensorart": "ta",
|
||||
"openmodeldb": "omdb",
|
||||
}
|
||||
|
||||
|
||||
@@ -288,6 +289,11 @@ class ModelSource:
|
||||
#: Sub-directory the "use default paths" template places downloads in.
|
||||
default_subdir: str = ""
|
||||
|
||||
#: Source id used to build the example URL shown in UI copy and error
|
||||
#: messages. ``owner/name`` suits repository sites; sites with a
|
||||
#: different identity shape override it with a real example.
|
||||
example_source_id: str = "user/repo"
|
||||
|
||||
#: 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
|
||||
@@ -340,6 +346,26 @@ class ModelSource:
|
||||
|
||||
raise NotImplementedError
|
||||
|
||||
def is_valid_source_id(self, source_id: str) -> bool:
|
||||
"""Return ``True`` when *source_id* is a safe id on this site.
|
||||
|
||||
Defaults to the ``owner/name`` repository rule; sites whose ids are
|
||||
not repositories (OpenModelDB's flat model ids) override it.
|
||||
"""
|
||||
|
||||
return is_valid_source_id(source_id)
|
||||
|
||||
def default_subdir_parts(self, source_id: str) -> tuple[str, ...]:
|
||||
"""Path segments appended to the model root by "use default paths".
|
||||
|
||||
Defaults to ``<default_subdir>/<owner>/<repo>`` so downloads from
|
||||
repository sites stay namespaced by author. Sites without an
|
||||
owner/repo split override it.
|
||||
"""
|
||||
|
||||
owner, repo_name = source_id.split("/", 1)
|
||||
return (self.default_subdir, owner, repo_name)
|
||||
|
||||
def asset_base_url(self, source_id: str, revision: str = "") -> str:
|
||||
"""Base URL used to resolve repository-relative asset paths."""
|
||||
|
||||
@@ -432,6 +458,18 @@ class ModelSource:
|
||||
f"{self.label or self.platform} does not support downloads", status=400
|
||||
)
|
||||
|
||||
async def resolve_download_url(
|
||||
self, source_id: str, filename: str, revision: str = ""
|
||||
) -> str:
|
||||
"""Resolve the download URL for one file, allowing async lookups.
|
||||
|
||||
Defaults to the synchronous :meth:`file_download_url`; sites whose
|
||||
download URL is not derivable from the id alone (OpenModelDB stores
|
||||
the URL inside its catalogue entry) override this to look it up.
|
||||
"""
|
||||
|
||||
return self.file_download_url(source_id, filename, revision)
|
||||
|
||||
def resolve_revision(self, revision: str = "") -> str:
|
||||
"""Return *revision*, falling back to this site's default branch."""
|
||||
|
||||
|
||||
@@ -0,0 +1,210 @@
|
||||
"""OpenModelDB model source (upscaler catalogue).
|
||||
|
||||
OpenModelDB (https://openmodeldb.info) is a static catalogue of upscaler
|
||||
models. Unlike the repository-based sources (Hugging Face, ModelScope) a
|
||||
model id here is a flat token (``4x-UltraSharp``) that *is* the published
|
||||
model identity: there is no owner/repo split, no revision, and no README.
|
||||
Everything the source needs — the resource download URLs, sizes, sha256
|
||||
hashes, tags and example images — comes from the site's bulk JSON dumps via
|
||||
:class:`~py.services.openmodeldb_client.OpenModelDBClient`, which caches the
|
||||
catalogue on disk, so every method below is a local lookup once warmed.
|
||||
|
||||
Only PyTorch resources (``.pth`` / ``.safetensors``) are listed for download:
|
||||
``.onnx`` is not a loadable weight format for the supported model types (see
|
||||
:data:`py.utils.constants.MODEL_FILE_EXTENSIONS`). Resources can carry
|
||||
mirror URLs; only the primary URL is ever used (see
|
||||
:meth:`OpenModelDBClient.primary_url`).
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import os
|
||||
import re
|
||||
from typing import Any, Optional
|
||||
from urllib.parse import urlparse
|
||||
|
||||
from .base import (
|
||||
ModelCardContext,
|
||||
ModelSource,
|
||||
ModelSourceCache,
|
||||
ModelSourceError,
|
||||
filter_weight_files,
|
||||
)
|
||||
from ..openmodeldb_client import OPENMODELDB_SITE_BASE, OpenModelDBClient
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
#: Model ids are flat tokens (``4x-UltraSharp``), usable as a path segment.
|
||||
_SOURCE_ID = re.compile(r"^[A-Za-z0-9_][A-Za-z0-9_.\-]*$")
|
||||
|
||||
_URL_PATTERN = re.compile(
|
||||
r"https?://(?:www\.)?openmodeldb\.info/models/(?P<id>[A-Za-z0-9_][A-Za-z0-9_.\-]*)"
|
||||
)
|
||||
_STRICT_URL_PATTERN = re.compile(
|
||||
r"https?://(?:www\.)?openmodeldb\.info/models/(?P<id>[A-Za-z0-9_][A-Za-z0-9_.\-]*)/?$"
|
||||
)
|
||||
|
||||
#: Resource platforms whose files ComfyUI can load.
|
||||
_DOWNLOADABLE_PLATFORMS = frozenset({"pytorch"})
|
||||
|
||||
|
||||
class OpenModelDBSource(ModelSource):
|
||||
"""OpenModelDB (``openmodeldb.info``)."""
|
||||
|
||||
platform = "openmodeldb"
|
||||
label = "OpenModelDB"
|
||||
supports_enrichment = True
|
||||
supports_download = True
|
||||
default_revision = ""
|
||||
default_subdir = "openmodeldb"
|
||||
example_source_id = "4x-UltraSharp"
|
||||
url_pattern = _URL_PATTERN
|
||||
strict_url_pattern = _STRICT_URL_PATTERN
|
||||
|
||||
def canonical_url(self, source_id: str) -> str:
|
||||
return f"{OPENMODELDB_SITE_BASE}/models/{source_id}"
|
||||
|
||||
def is_valid_source_id(self, source_id: str) -> bool:
|
||||
"""OpenModelDB ids are flat tokens, not ``owner/name`` repositories."""
|
||||
|
||||
return bool(isinstance(source_id, str) and _SOURCE_ID.match(source_id))
|
||||
|
||||
def default_subdir_parts(self, source_id: str) -> tuple[str, ...]:
|
||||
"""Flat catalogue: there is no owner/repo split to mirror on disk."""
|
||||
|
||||
return (self.default_subdir,)
|
||||
|
||||
async def fetch_model_card_context(
|
||||
self,
|
||||
source_id: str,
|
||||
filename: str = "",
|
||||
*,
|
||||
sha256: str = "",
|
||||
cache: Optional["ModelSourceCache"] = None,
|
||||
) -> ModelCardContext:
|
||||
"""Build the card extras from the cached catalogue entry.
|
||||
|
||||
OpenModelDB has no README; the catalogue entry itself carries the
|
||||
description, license, tags and example images, so the context is the
|
||||
whole card. The catalogue is bulk-loaded and disk-cached, so no
|
||||
per-run memo is needed.
|
||||
"""
|
||||
|
||||
try:
|
||||
client = await OpenModelDBClient.get_instance()
|
||||
found = await client.get_model_entry(source_id)
|
||||
except Exception as exc: # never break enrichment on a lookup fault
|
||||
logger.debug("OpenModelDB context lookup failed for %s: %s", source_id, exc)
|
||||
return ModelCardContext()
|
||||
if found is None:
|
||||
return ModelCardContext()
|
||||
entry = found[1]
|
||||
|
||||
scale = entry.get("scale")
|
||||
arch_name = client._resolve_architecture_name(entry)
|
||||
# e.g. "ESRGAN 4x" — closest thing upscalers have to a base model,
|
||||
# recorded as a hint rather than a canonical base-model name.
|
||||
base_hint = (
|
||||
f"{arch_name} {scale}x".strip()
|
||||
if arch_name and isinstance(scale, (int, float))
|
||||
else arch_name
|
||||
)
|
||||
description = entry.get("description")
|
||||
|
||||
return ModelCardContext(
|
||||
description=description if isinstance(description, str) else "",
|
||||
model_name=entry.get("name") or source_id,
|
||||
license=entry.get("license") or "",
|
||||
model_type="Upscaler",
|
||||
base_model_aliases=[base_hint] if base_hint else [],
|
||||
official_tags=client._resolve_tags(entry),
|
||||
example_images=client.example_image_urls(entry),
|
||||
source_model_id=source_id,
|
||||
)
|
||||
|
||||
async def list_files(
|
||||
self, source_id: str, revision: str = ""
|
||||
) -> list[dict[str, Any]]:
|
||||
"""List the entry's directly downloadable PyTorch resources.
|
||||
|
||||
Resources whose only mirrors are HTML-gateway hosts (mediafire,
|
||||
mega.nz, drive.google.com) are skipped: they serve a web page, not
|
||||
the file bytes. When every resource is mirror-only this raises a
|
||||
manual-download hint instead of returning an empty list, which the
|
||||
download dialog would otherwise misreport as "no model files".
|
||||
"""
|
||||
|
||||
client = await OpenModelDBClient.get_instance()
|
||||
entry = await self._require_entry(client, source_id)
|
||||
|
||||
saw_mirror_only = False
|
||||
entries = []
|
||||
for resource in entry.get("resources") or []:
|
||||
if not isinstance(resource, dict):
|
||||
continue
|
||||
if str(resource.get("platform") or "").lower() not in _DOWNLOADABLE_PLATFORMS:
|
||||
continue
|
||||
url = client.direct_url(resource)
|
||||
if not url:
|
||||
saw_mirror_only = True
|
||||
continue
|
||||
size = resource.get("size")
|
||||
entries.append(
|
||||
(
|
||||
client.resource_filename(source_id, resource),
|
||||
size if isinstance(size, (int, float)) else 0,
|
||||
)
|
||||
)
|
||||
|
||||
if not entries and saw_mirror_only:
|
||||
raise ModelSourceError(
|
||||
f"None of this model's mirrors support direct download; "
|
||||
f"download it manually from {self.canonical_url(source_id)}",
|
||||
status=400,
|
||||
)
|
||||
|
||||
return filter_weight_files(entries)
|
||||
|
||||
async def resolve_download_url(
|
||||
self, source_id: str, filename: str, revision: str = ""
|
||||
) -> str:
|
||||
"""Resolve the direct download URL of one resource by filename."""
|
||||
|
||||
client = await OpenModelDBClient.get_instance()
|
||||
entry = await self._require_entry(client, source_id)
|
||||
|
||||
resource = client.find_resource_by_filename(
|
||||
source_id, entry, os.path.basename(filename)
|
||||
)
|
||||
if resource is None:
|
||||
raise ModelSourceError(
|
||||
f"'{filename}' is not a downloadable resource of '{source_id}'",
|
||||
status=404,
|
||||
)
|
||||
url = client.direct_url(resource)
|
||||
if not url:
|
||||
host = urlparse(client.primary_url(resource)).netloc or "this mirror"
|
||||
raise ModelSourceError(
|
||||
f"This mirror ({host}) requires manual download from "
|
||||
f"{self.canonical_url(source_id)}",
|
||||
status=400,
|
||||
)
|
||||
return url
|
||||
|
||||
async def _require_entry(
|
||||
self, client: OpenModelDBClient, source_id: str
|
||||
) -> dict[str, Any]:
|
||||
"""Return the catalogue entry, raising a mapped error otherwise."""
|
||||
|
||||
if not await client.catalogue_ready():
|
||||
raise ModelSourceError("OpenModelDB catalogue unavailable", status=502)
|
||||
found = await client.get_model_entry(source_id)
|
||||
if found is None:
|
||||
raise ModelSourceError(
|
||||
f"Model '{source_id}' not found on OpenModelDB", status=404
|
||||
)
|
||||
return found[1]
|
||||
|
||||
|
||||
__all__ = ["OpenModelDBSource"]
|
||||
@@ -14,6 +14,7 @@ from typing import Any, Dict, Mapping, Optional
|
||||
from .base import GROUP_PREFIXES, ModelSource, SourceRef, clean_source_url
|
||||
from .huggingface import HuggingFaceSource
|
||||
from .modelscope import ModelScopeIntlSource, ModelScopeSource
|
||||
from .openmodeldb import OpenModelDBSource
|
||||
from .tensorart import TensorArtSource
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
@@ -26,6 +27,7 @@ _SOURCES: tuple[ModelSource, ...] = (
|
||||
ModelScopeSource(),
|
||||
ModelScopeIntlSource(),
|
||||
TensorArtSource(),
|
||||
OpenModelDBSource(),
|
||||
)
|
||||
|
||||
_BY_PLATFORM: Dict[str, ModelSource] = {s.platform: s for s in _SOURCES}
|
||||
|
||||
@@ -42,6 +42,7 @@ class TensorArtSource(ModelSource):
|
||||
label = "TensorArt"
|
||||
supports_enrichment = False
|
||||
supports_download = False
|
||||
example_source_id = "827823520299086029"
|
||||
url_pattern = _URL_PATTERN
|
||||
strict_url_pattern = _STRICT_URL_PATTERN
|
||||
|
||||
|
||||
@@ -0,0 +1,762 @@
|
||||
"""Client for the OpenModelDB bulk JSON API.
|
||||
|
||||
OpenModelDB (https://openmodeldb.info) is a static catalogue of upscaler
|
||||
models. It exposes no per-model or by-hash endpoint — only bulk JSON dumps
|
||||
(``/api/v1/models.json`` and friends), so this client downloads the dumps
|
||||
once, caches them on disk with a TTL, honors ETag/Last-Modified on refresh,
|
||||
and builds an in-memory SHA256 -> model index for read-only metadata lookups.
|
||||
|
||||
Every catalogue resource carries a ``sha256`` and a byte ``size``, which is
|
||||
what makes hash-based matching against local files possible. Lookups degrade
|
||||
gracefully: when the catalogue cannot be fetched (offline, upstream failure)
|
||||
the stale disk cache is used, and if there is no cache at all the lookup
|
||||
reports "not found" so the metadata fallback chain simply moves on.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import json
|
||||
import logging
|
||||
import os
|
||||
import time
|
||||
from typing import Any, Dict, List, Optional, Tuple
|
||||
from urllib.parse import urlparse
|
||||
|
||||
from .downloader import get_downloader
|
||||
from .errors import RateLimitError
|
||||
from ..utils.cache_paths import get_cache_base_dir
|
||||
from ..utils.constants import MODEL_FILE_EXTENSIONS
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
OPENMODELDB_API_BASE = "https://openmodeldb.info/api/v1"
|
||||
OPENMODELDB_SITE_BASE = "https://openmodeldb.info"
|
||||
|
||||
#: Value emitted as ``source`` in the synthesized version dict; the metadata
|
||||
#: sync service persists it as the model's ``metadata_source``.
|
||||
METADATA_SOURCE_VALUE = "openmodeldb"
|
||||
|
||||
#: Hosts whose URLs serve an HTML interstitial page instead of the raw file
|
||||
#: bytes. Downloading from them would silently save a web page as ``.pth``
|
||||
#: (the download flow does not verify the sha256 afterwards), so the model
|
||||
#: source layer rejects them with a manual-download hint. Kept deliberately
|
||||
#: small and explicit.
|
||||
HTML_GATEWAY_HOSTS = frozenset({"mediafire.com", "mega.nz", "drive.google.com"})
|
||||
|
||||
|
||||
def is_html_gateway_url(url: str) -> bool:
|
||||
"""Return ``True`` when *url* points at a known HTML-gateway host."""
|
||||
|
||||
if not isinstance(url, str) or not url:
|
||||
return False
|
||||
try:
|
||||
host = urlparse(url).netloc.lower()
|
||||
except ValueError:
|
||||
return False
|
||||
return any(host == g or host.endswith(f".{g}") for g in HTML_GATEWAY_HOSTS)
|
||||
|
||||
|
||||
def _is_ephemeral_viewer_url(url: str) -> bool:
|
||||
"""Return ``True`` for imgdiff.net session URLs.
|
||||
|
||||
Paired comparisons are hosted as ephemeral imgdiff viewer sessions
|
||||
(``/api/image.php?id=...``) that expire shortly after the site build;
|
||||
they 404 when used as an ``<img>`` source and must never be emitted as a
|
||||
displayable image URL.
|
||||
"""
|
||||
|
||||
return isinstance(url, str) and "imgdiff.net/api/" in url
|
||||
|
||||
#: Bulk dumps consumed by the client. Only ``models`` is strictly required;
|
||||
#: the rest resolve ids to human-readable names and degrade to raw ids.
|
||||
_DUMP_NAMES = ("models", "users", "tags", "architectures")
|
||||
|
||||
#: How long a fetched catalogue is considered fresh before a revalidation
|
||||
#: request is made. The site only changes when it is rebuilt (hours to days),
|
||||
#: so a daily TTL avoids re-downloading the ~1.4MB models dump on every
|
||||
#: lookup while still picking up new models reasonably fast.
|
||||
CACHE_TTL_SECONDS = 24 * 60 * 60
|
||||
|
||||
_META_FILENAME = "_meta.json"
|
||||
|
||||
|
||||
class OpenModelDBClient:
|
||||
"""Hash-lookup client over a locally cached OpenModelDB catalogue dump."""
|
||||
|
||||
_instance: Optional["OpenModelDBClient"] = None
|
||||
_instance_lock = asyncio.Lock()
|
||||
|
||||
@classmethod
|
||||
async def get_instance(cls) -> "OpenModelDBClient":
|
||||
"""Get the singleton instance of OpenModelDBClient."""
|
||||
async with cls._instance_lock:
|
||||
if cls._instance is None:
|
||||
cls._instance = cls()
|
||||
|
||||
# Register this client as a metadata provider (mirrors the
|
||||
# CivitAI/CivArchive client bootstrap).
|
||||
from .model_metadata_provider import (
|
||||
ModelMetadataProviderManager,
|
||||
OpenModelDBModelMetadataProvider,
|
||||
)
|
||||
|
||||
provider_manager = await ModelMetadataProviderManager.get_instance()
|
||||
provider_manager.register_provider(
|
||||
"openmodeldb",
|
||||
OpenModelDBModelMetadataProvider(cls._instance),
|
||||
False,
|
||||
)
|
||||
|
||||
return cls._instance
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
cache_dir: Optional[str] = None,
|
||||
ttl_seconds: float = CACHE_TTL_SECONDS,
|
||||
) -> None:
|
||||
# Guard re-initialization for the singleton pattern.
|
||||
if hasattr(self, "_initialized"):
|
||||
return
|
||||
self._initialized = True
|
||||
|
||||
self._cache_dir_override = cache_dir
|
||||
self._ttl_seconds = ttl_seconds
|
||||
|
||||
self._models: Dict[str, Dict[str, Any]] = {}
|
||||
self._users: Dict[str, Dict[str, Any]] = {}
|
||||
self._tags: Dict[str, Dict[str, Any]] = {}
|
||||
self._architectures: Dict[str, Dict[str, Any]] = {}
|
||||
# sha256 (lowercase) -> (model_id, model entry, matching resource)
|
||||
self._index: Dict[str, Tuple[str, Dict[str, Any], Dict[str, Any]]] = {}
|
||||
self._loaded_at: float = 0.0
|
||||
self._load_lock = asyncio.Lock()
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# Public API
|
||||
# ------------------------------------------------------------------
|
||||
|
||||
async def get_model_by_hash(
|
||||
self, model_hash: str
|
||||
) -> Tuple[Optional[Dict[str, Any]], Optional[str]]:
|
||||
"""Find an upscaler model by SHA256 hash.
|
||||
|
||||
Returns a CivitAI-shaped version dict (same contract as the other
|
||||
metadata providers) or ``(None, reason)``.
|
||||
"""
|
||||
if not model_hash or not isinstance(model_hash, str):
|
||||
return None, "Model not found"
|
||||
|
||||
try:
|
||||
loaded = await self._ensure_loaded()
|
||||
except RateLimitError:
|
||||
raise
|
||||
except Exception as exc:
|
||||
logger.error("OpenModelDB lookup failed for %s: %s", model_hash[:10], exc)
|
||||
return None, str(exc)
|
||||
|
||||
if not loaded:
|
||||
return None, "OpenModelDB catalogue unavailable"
|
||||
|
||||
hit = self._index.get(model_hash.lower())
|
||||
if hit is None:
|
||||
return None, "Model not found"
|
||||
|
||||
model_id, model_entry, resource = hit
|
||||
return self._to_civitai_version(model_id, model_entry, resource), None
|
||||
|
||||
async def catalogue_ready(self) -> bool:
|
||||
"""Return ``True`` when the catalogue is loaded (or loadable)."""
|
||||
|
||||
try:
|
||||
return await self._ensure_loaded()
|
||||
except RateLimitError:
|
||||
raise
|
||||
except Exception as exc:
|
||||
logger.error("OpenModelDB catalogue load failed: %s", exc)
|
||||
return False
|
||||
|
||||
async def get_model_entry(
|
||||
self, model_id: str
|
||||
) -> Optional[Tuple[str, Dict[str, Any]]]:
|
||||
"""Return ``(model_id, catalogue entry)`` for *model_id*, or ``None``.
|
||||
|
||||
Loads the catalogue on first use; an unavailable catalogue and an
|
||||
unknown id both yield ``None`` (callers that need to distinguish the
|
||||
two can check :meth:`catalogue_ready` first).
|
||||
"""
|
||||
if not model_id or not isinstance(model_id, str):
|
||||
return None
|
||||
if not await self.catalogue_ready():
|
||||
return None
|
||||
entry = self._models.get(model_id)
|
||||
if not isinstance(entry, dict):
|
||||
return None
|
||||
return model_id, entry
|
||||
|
||||
def find_resource_by_filename(
|
||||
self, model_id: str, entry: Dict[str, Any], filename: str
|
||||
) -> Optional[Dict[str, Any]]:
|
||||
"""Match a resource by its derived filename (see :meth:`resource_filename`)."""
|
||||
target = (filename or "").strip().lower()
|
||||
if not target:
|
||||
return None
|
||||
for resource in entry.get("resources") or []:
|
||||
if not isinstance(resource, dict):
|
||||
continue
|
||||
if self.resource_filename(model_id, resource).lower() == target:
|
||||
return resource
|
||||
return None
|
||||
|
||||
@staticmethod
|
||||
def resource_filename(model_id: str, resource: Dict[str, Any]) -> str:
|
||||
"""Derive the local filename for a catalogue resource.
|
||||
|
||||
The download URL's basename is not authoritative — mirrors like
|
||||
mediafire put the real filename mid-path
|
||||
(``/file/<key>/90s_Sonic_2x.pth/file``) and folder links (mega.nz)
|
||||
have no filename at all. Strategy: the first URL path segment whose
|
||||
extension is a known model format, else ``{model_id}.{type}`` (the
|
||||
catalogue's ``type`` field is authoritative).
|
||||
"""
|
||||
urls = resource.get("urls")
|
||||
for url in urls if isinstance(urls, list) else []:
|
||||
if not isinstance(url, str):
|
||||
continue
|
||||
path = url.split("?", 1)[0].split("#", 1)[0]
|
||||
for segment in path.split("/"):
|
||||
if os.path.splitext(segment)[1].lower() in MODEL_FILE_EXTENSIONS:
|
||||
return segment
|
||||
resource_type = str(resource.get("type") or "").lower()
|
||||
extension = (
|
||||
resource_type if resource_type in {"pth", "safetensors", "onnx"} else "bin"
|
||||
)
|
||||
return f"{model_id}.{extension}"
|
||||
|
||||
@staticmethod
|
||||
def primary_url(resource: Dict[str, Any]) -> str:
|
||||
"""Return the resource's primary download URL.
|
||||
|
||||
Only the first URL is used: additional entries are mirrors that may
|
||||
need site-specific handling (e.g. mega.nz) and are never tried
|
||||
automatically.
|
||||
"""
|
||||
urls = resource.get("urls")
|
||||
if isinstance(urls, list):
|
||||
for url in urls:
|
||||
if isinstance(url, str) and url.startswith("http"):
|
||||
return url
|
||||
return ""
|
||||
|
||||
@staticmethod
|
||||
def direct_url(resource: Dict[str, Any]) -> str:
|
||||
"""Return the first URL that serves raw bytes, or ``""``.
|
||||
|
||||
HTML-gateway hosts (mediafire, mega.nz, drive.google.com — see
|
||||
:data:`HTML_GATEWAY_HOSTS`) serve an interstitial page instead of the
|
||||
file, so they are skipped here and reported to the user instead.
|
||||
"""
|
||||
urls = resource.get("urls")
|
||||
if isinstance(urls, list):
|
||||
for url in urls:
|
||||
if (
|
||||
isinstance(url, str)
|
||||
and url.startswith("http")
|
||||
and not is_html_gateway_url(url)
|
||||
):
|
||||
return url
|
||||
return ""
|
||||
|
||||
@staticmethod
|
||||
def _absolutize(url: str) -> str:
|
||||
"""Turn a site-relative path (``/thumbs/...``) into an absolute URL."""
|
||||
|
||||
if isinstance(url, str) and url.startswith("/"):
|
||||
return f"{OPENMODELDB_SITE_BASE}{url}"
|
||||
return url
|
||||
|
||||
def _paired_display_url(self, image: Dict[str, Any]) -> str:
|
||||
"""Return the displayable URL for a paired comparison image.
|
||||
|
||||
Prefers the site-hosted thumbnail: the ``LR``/``SR`` originals are
|
||||
frequently ephemeral imgdiff session URLs that 404 outside the
|
||||
viewer. Falls back to the SR (then LR) original only when it is not
|
||||
one of those session URLs.
|
||||
"""
|
||||
thumbnail = image.get("thumbnail")
|
||||
if isinstance(thumbnail, str) and thumbnail:
|
||||
return self._absolutize(thumbnail)
|
||||
for key in ("SR", "LR"):
|
||||
original = image.get(key)
|
||||
if (
|
||||
isinstance(original, str)
|
||||
and original
|
||||
and not _is_ephemeral_viewer_url(original)
|
||||
):
|
||||
return original
|
||||
return ""
|
||||
|
||||
def _model_preview_url(self, entry: Dict[str, Any]) -> str:
|
||||
"""Return the model-level thumbnail URL, mirroring the site's own
|
||||
``getPreviewImage`` precedence (paired → SR, standalone → url)."""
|
||||
thumbnail = entry.get("thumbnail")
|
||||
if not isinstance(thumbnail, dict):
|
||||
return ""
|
||||
if thumbnail.get("type") == "paired":
|
||||
url = thumbnail.get("SR") or thumbnail.get("LR")
|
||||
else:
|
||||
url = thumbnail.get("url")
|
||||
if isinstance(url, str) and url:
|
||||
return self._absolutize(url)
|
||||
return ""
|
||||
|
||||
def example_image_urls(self, entry: Dict[str, Any]) -> List[str]:
|
||||
"""Return displayable example-image URLs, model thumbnail first.
|
||||
|
||||
Never contains ephemeral imgdiff session URLs; standalone images keep
|
||||
their direct URLs (regular image hosts are hotlinkable).
|
||||
"""
|
||||
urls: List[str] = []
|
||||
lead = self._model_preview_url(entry)
|
||||
if lead:
|
||||
urls.append(lead)
|
||||
for image in entry.get("images") or []:
|
||||
if not isinstance(image, dict):
|
||||
continue
|
||||
if image.get("type") == "paired":
|
||||
url = self._paired_display_url(image)
|
||||
else:
|
||||
url = image.get("url")
|
||||
if isinstance(url, str) and url and url not in urls:
|
||||
urls.append(url)
|
||||
return urls
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# Catalogue loading
|
||||
# ------------------------------------------------------------------
|
||||
|
||||
async def _ensure_loaded(self) -> bool:
|
||||
"""Ensure the in-memory index is built, refreshing stale caches."""
|
||||
async with self._load_lock:
|
||||
if self._index and (time.monotonic() - self._loaded_at) < self._ttl_seconds:
|
||||
return True
|
||||
|
||||
meta = self._read_meta()
|
||||
fetched_at = float(meta.get("fetched_at") or 0.0)
|
||||
disk_fresh = (
|
||||
fetched_at > 0
|
||||
and (time.time() - fetched_at) < self._ttl_seconds
|
||||
and all(os.path.exists(self._dump_path(name)) for name in _DUMP_NAMES)
|
||||
)
|
||||
|
||||
if disk_fresh:
|
||||
if self._load_from_disk():
|
||||
return True
|
||||
# Corrupt disk cache: fall through to a network refresh.
|
||||
|
||||
if await self._refresh_from_network(meta):
|
||||
return True
|
||||
|
||||
# Network failed or was blocked: fall back to whatever is on disk,
|
||||
# however stale — old metadata beats none.
|
||||
if fetched_at > 0 and self._load_from_disk():
|
||||
logger.info("Using stale OpenModelDB cache (network refresh failed)")
|
||||
return True
|
||||
|
||||
return False
|
||||
|
||||
def _load_from_disk(self) -> bool:
|
||||
"""Load all dumps from the disk cache and rebuild the index."""
|
||||
payloads: Dict[str, Dict[str, Any]] = {}
|
||||
for name in _DUMP_NAMES:
|
||||
path = self._dump_path(name)
|
||||
try:
|
||||
with open(path, "r", encoding="utf-8") as handle:
|
||||
data = json.load(handle)
|
||||
except FileNotFoundError:
|
||||
if name == "models":
|
||||
return False
|
||||
data = {}
|
||||
except (OSError, json.JSONDecodeError) as exc:
|
||||
logger.warning("Failed to read OpenModelDB cache %s: %s", path, exc)
|
||||
if name == "models":
|
||||
return False
|
||||
data = {}
|
||||
payloads[name] = data if isinstance(data, dict) else {}
|
||||
|
||||
if not payloads["models"]:
|
||||
return False
|
||||
|
||||
self._install_payloads(payloads)
|
||||
return True
|
||||
|
||||
async def _refresh_from_network(self, meta: Dict[str, Any]) -> bool:
|
||||
"""Revalidate cached dumps against the site and rebuild the index.
|
||||
|
||||
Honors ETag/Last-Modified via a HEAD probe: an unchanged dump keeps
|
||||
its cached body, so a TTL expiry without upstream changes costs one
|
||||
tiny request per dump instead of a full download.
|
||||
"""
|
||||
payloads: Dict[str, Dict[str, Any]] = {}
|
||||
etags: Dict[str, str] = dict(meta.get("etags") or {})
|
||||
last_modified: Dict[str, str] = dict(meta.get("last_modified") or {})
|
||||
|
||||
for name in _DUMP_NAMES:
|
||||
payload, etag, modified = await self._fetch_dump(
|
||||
name,
|
||||
known_etag=etags.get(name) or "",
|
||||
known_last_modified=last_modified.get(name) or "",
|
||||
)
|
||||
if payload is None:
|
||||
if name == "models":
|
||||
return False
|
||||
payload = {}
|
||||
payloads[name] = payload
|
||||
if etag:
|
||||
etags[name] = etag
|
||||
if modified:
|
||||
last_modified[name] = modified
|
||||
|
||||
self._install_payloads(payloads)
|
||||
self._write_cache(payloads, etags, last_modified)
|
||||
return True
|
||||
|
||||
async def _fetch_dump(
|
||||
self,
|
||||
name: str,
|
||||
*,
|
||||
known_etag: str,
|
||||
known_last_modified: str,
|
||||
) -> Tuple[Optional[Dict[str, Any]], str, str]:
|
||||
"""Fetch one dump, returning ``(body, etag, last_modified)``.
|
||||
|
||||
``body`` is ``None`` when the fetch failed and there is no usable
|
||||
cached copy. When the HEAD probe shows the resource unchanged, the
|
||||
cached body is returned without a full download.
|
||||
"""
|
||||
url = f"{OPENMODELDB_API_BASE}/{name}.json"
|
||||
disk_path = self._dump_path(name)
|
||||
have_cached = os.path.exists(disk_path)
|
||||
|
||||
downloader = await get_downloader()
|
||||
|
||||
head_etag = ""
|
||||
head_modified = ""
|
||||
try:
|
||||
head_ok, head_headers = await downloader.get_response_headers(url)
|
||||
except Exception as exc: # pragma: no cover - defensive guard
|
||||
logger.debug("OpenModelDB HEAD probe failed for %s: %s", url, exc)
|
||||
head_ok, head_headers = False, {}
|
||||
|
||||
if head_ok and isinstance(head_headers, dict):
|
||||
# aiohttp headers are case-insensitive; plain dicts in tests are not.
|
||||
head_etag = str(head_headers.get("ETag") or head_headers.get("etag") or "")
|
||||
head_modified = str(
|
||||
head_headers.get("Last-Modified") or head_headers.get("last-modified") or ""
|
||||
)
|
||||
|
||||
if (
|
||||
have_cached
|
||||
and known_etag
|
||||
and head_etag
|
||||
and head_etag == known_etag
|
||||
):
|
||||
cached = self._read_dump_file(disk_path)
|
||||
if cached is not None:
|
||||
logger.debug("OpenModelDB %s unchanged (etag match); using cache", name)
|
||||
return cached, known_etag, known_last_modified or head_modified
|
||||
|
||||
success, payload = await downloader.make_request("GET", url, use_auth=False)
|
||||
if isinstance(payload, RateLimitError):
|
||||
raise payload
|
||||
if not success or not isinstance(payload, dict):
|
||||
logger.warning(
|
||||
"OpenModelDB %s fetch failed: %s",
|
||||
name,
|
||||
payload if isinstance(payload, str) else "unexpected payload",
|
||||
)
|
||||
if have_cached:
|
||||
cached = self._read_dump_file(disk_path)
|
||||
if cached is not None:
|
||||
return cached, known_etag, known_last_modified
|
||||
return None, known_etag, known_last_modified
|
||||
|
||||
return payload, head_etag or known_etag, head_modified or known_last_modified
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# Index and transformation
|
||||
# ------------------------------------------------------------------
|
||||
|
||||
def _install_payloads(self, payloads: Dict[str, Dict[str, Any]]) -> None:
|
||||
"""Install dump payloads and rebuild the sha256 index."""
|
||||
self._models = payloads.get("models") or {}
|
||||
self._users = payloads.get("users") or {}
|
||||
self._tags = payloads.get("tags") or {}
|
||||
self._architectures = payloads.get("architectures") or {}
|
||||
self._index = self._build_index(self._models)
|
||||
self._loaded_at = time.monotonic()
|
||||
logger.debug(
|
||||
"OpenModelDB catalogue loaded: %d models, %d indexed hashes",
|
||||
len(self._models),
|
||||
len(self._index),
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _build_index(
|
||||
models: Dict[str, Dict[str, Any]]
|
||||
) -> Dict[str, Tuple[str, Dict[str, Any], Dict[str, Any]]]:
|
||||
"""Build the sha256 -> (model_id, model entry, resource) index."""
|
||||
index: Dict[str, Tuple[str, Dict[str, Any], Dict[str, Any]]] = {}
|
||||
for model_id, entry in models.items():
|
||||
if not isinstance(entry, dict):
|
||||
continue
|
||||
resources = entry.get("resources")
|
||||
if not isinstance(resources, list):
|
||||
continue
|
||||
for resource in resources:
|
||||
if not isinstance(resource, dict):
|
||||
continue
|
||||
sha256 = resource.get("sha256")
|
||||
if not isinstance(sha256, str) or not sha256:
|
||||
continue
|
||||
# First writer wins: duplicate hashes across catalogue entries
|
||||
# are ambiguous and cannot be disambiguated locally.
|
||||
index.setdefault(sha256.lower(), (model_id, entry, resource))
|
||||
return index
|
||||
|
||||
def _resolve_authors(self, entry: Dict[str, Any]) -> Tuple[str, List[str]]:
|
||||
"""Resolve the author field to a display name plus the raw user ids."""
|
||||
raw = entry.get("author")
|
||||
author_ids = raw if isinstance(raw, list) else [raw]
|
||||
ids = [str(a) for a in author_ids if isinstance(a, str) and a]
|
||||
names: List[str] = []
|
||||
for author_id in ids:
|
||||
user = self._users.get(author_id)
|
||||
name = user.get("name") if isinstance(user, dict) else None
|
||||
names.append(name if isinstance(name, str) and name else author_id)
|
||||
return ", ".join(names), ids
|
||||
|
||||
def _resolve_tags(self, entry: Dict[str, Any]) -> List[str]:
|
||||
"""Resolve tag ids to their display names."""
|
||||
raw_tags = entry.get("tags")
|
||||
if not isinstance(raw_tags, list):
|
||||
return []
|
||||
resolved: List[str] = []
|
||||
for tag_id in raw_tags:
|
||||
if not isinstance(tag_id, str) or not tag_id:
|
||||
continue
|
||||
tag = self._tags.get(tag_id)
|
||||
name = tag.get("name") if isinstance(tag, dict) else None
|
||||
resolved.append(name if isinstance(name, str) and name else tag_id)
|
||||
return resolved
|
||||
|
||||
def _resolve_architecture_name(self, entry: Dict[str, Any]) -> str:
|
||||
"""Resolve the architecture id to its display name."""
|
||||
arch_id = entry.get("architecture")
|
||||
if not isinstance(arch_id, str) or not arch_id:
|
||||
return ""
|
||||
arch = self._architectures.get(arch_id)
|
||||
if isinstance(arch, dict):
|
||||
name = arch.get("name")
|
||||
if isinstance(name, str) and name:
|
||||
return name
|
||||
return arch_id
|
||||
|
||||
@staticmethod
|
||||
def _resource_format(resource: Dict[str, Any]) -> str:
|
||||
"""Map an OpenModelDB resource type to a CivitAI file metadata format."""
|
||||
resource_type = str(resource.get("type") or "").lower()
|
||||
if resource_type == "safetensors":
|
||||
return "SafeTensor"
|
||||
if resource_type in ("pth", "pt", "ckpt"):
|
||||
return "PickleTensor"
|
||||
return "Other"
|
||||
|
||||
def _to_civitai_version(
|
||||
self,
|
||||
model_id: str,
|
||||
entry: Dict[str, Any],
|
||||
matched_resource: Dict[str, Any],
|
||||
) -> Dict[str, Any]:
|
||||
"""Map an OpenModelDB catalogue entry to a CivitAI-shaped version dict.
|
||||
|
||||
Follows the same contract as the CivArchive/SQLite providers so the
|
||||
metadata sync service can merge it unchanged. Numeric ``id``/``modelId``
|
||||
are deliberately omitted: OpenModelDB ids are strings, and consumers
|
||||
treat a missing ``modelId`` as "not a CivitAI model" (no CivitAI page
|
||||
link, no update checks).
|
||||
"""
|
||||
author_display, author_ids = self._resolve_authors(entry)
|
||||
tags = self._resolve_tags(entry)
|
||||
architecture_id = entry.get("architecture")
|
||||
architecture_name = self._resolve_architecture_name(entry)
|
||||
description = entry.get("description")
|
||||
license_name = entry.get("license")
|
||||
page_url = f"{OPENMODELDB_SITE_BASE}/models/{model_id}"
|
||||
|
||||
files: List[Dict[str, Any]] = []
|
||||
resources = entry.get("resources")
|
||||
for resource in resources if isinstance(resources, list) else []:
|
||||
if not isinstance(resource, dict):
|
||||
continue
|
||||
# The displayable download URL prefers a direct-bytes mirror when
|
||||
# one exists; the filename is derived (never the raw URL basename,
|
||||
# which mediafire-style mirrors leave as "file").
|
||||
download_url = self.direct_url(resource) or self.primary_url(resource)
|
||||
sha256 = resource.get("sha256")
|
||||
size_bytes = resource.get("size")
|
||||
files.append(
|
||||
{
|
||||
"name": self.resource_filename(model_id, resource),
|
||||
"type": "Model",
|
||||
"sizeKB": (size_bytes / 1024.0)
|
||||
if isinstance(size_bytes, (int, float))
|
||||
else 0,
|
||||
"downloadUrl": download_url,
|
||||
"primary": resource is matched_resource,
|
||||
"hashes": {"SHA256": str(sha256).upper()} if sha256 else {},
|
||||
"metadata": {"format": self._resource_format(resource)},
|
||||
}
|
||||
)
|
||||
|
||||
images: List[Dict[str, Any]] = []
|
||||
# The model-level thumbnail is the site's own preview pick and larger
|
||||
# than the per-image small thumbs; the card preview derives from
|
||||
# images[0], so it leads the list.
|
||||
lead = self._model_preview_url(entry)
|
||||
if lead:
|
||||
images.append({"url": lead, "nsfwLevel": 1, "type": "image"})
|
||||
raw_images = entry.get("images")
|
||||
for image in raw_images if isinstance(raw_images, list) else []:
|
||||
if not isinstance(image, dict):
|
||||
continue
|
||||
paired = image.get("type") == "paired"
|
||||
# Paired entries show the upscaled (SR) result as the preview.
|
||||
url = self._paired_display_url(image) if paired else image.get("url")
|
||||
if not isinstance(url, str) or not url:
|
||||
continue
|
||||
if any(existing["url"] == url for existing in images):
|
||||
continue
|
||||
mapped: Dict[str, Any] = {"url": url, "nsfwLevel": 1, "type": "image"}
|
||||
thumbnail = image.get("thumbnail")
|
||||
if isinstance(thumbnail, str) and thumbnail:
|
||||
thumbnail_url = self._absolutize(thumbnail)
|
||||
if thumbnail_url != url:
|
||||
mapped["thumbnailUrl"] = thumbnail_url
|
||||
meta: Dict[str, Any] = {}
|
||||
if paired:
|
||||
comparison = image.get("SR") or image.get("LR")
|
||||
if (
|
||||
isinstance(comparison, str)
|
||||
and comparison
|
||||
and comparison != url
|
||||
and _is_ephemeral_viewer_url(comparison)
|
||||
):
|
||||
# Ephemeral imgdiff viewer session, kept for reference
|
||||
# only — it 404s outside the session and is never
|
||||
# displayable.
|
||||
meta["comparisonUrl"] = comparison
|
||||
caption = image.get("caption")
|
||||
if isinstance(caption, str) and caption:
|
||||
meta["caption"] = caption
|
||||
if meta:
|
||||
mapped["meta"] = meta
|
||||
images.append(mapped)
|
||||
|
||||
return {
|
||||
"name": entry.get("name") or model_id,
|
||||
# Upscalers are not tied to a diffusion base model.
|
||||
"baseModel": "Other",
|
||||
"description": description or "",
|
||||
"publishedAt": entry.get("date"),
|
||||
"trainedWords": [],
|
||||
"model": {
|
||||
"name": entry.get("name") or model_id,
|
||||
"type": "Upscaler",
|
||||
"nsfw": False,
|
||||
"description": description,
|
||||
"tags": tags,
|
||||
"license": license_name or "",
|
||||
},
|
||||
"creator": {"username": author_display, "image": None},
|
||||
"files": files,
|
||||
"images": images,
|
||||
"source": METADATA_SOURCE_VALUE,
|
||||
# OpenModelDB-native provenance, kept inside the persisted payload
|
||||
# so the UI can link to the model page in a later phase.
|
||||
"openmodeldb": {
|
||||
"id": model_id,
|
||||
"url": page_url,
|
||||
"authors": author_ids,
|
||||
"architecture": architecture_id or "",
|
||||
"architectureName": architecture_name,
|
||||
"scale": entry.get("scale"),
|
||||
"inputChannels": entry.get("inputChannels"),
|
||||
"outputChannels": entry.get("outputChannels"),
|
||||
"size": entry.get("size") or [],
|
||||
"license": license_name or "",
|
||||
"date": entry.get("date"),
|
||||
},
|
||||
}
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# Disk cache
|
||||
# ------------------------------------------------------------------
|
||||
|
||||
def _cache_dir(self) -> str:
|
||||
base = self._cache_dir_override or os.path.join(
|
||||
get_cache_base_dir(), "openmodeldb"
|
||||
)
|
||||
os.makedirs(base, exist_ok=True)
|
||||
return base
|
||||
|
||||
def _dump_path(self, name: str) -> str:
|
||||
return os.path.join(self._cache_dir(), f"{name}.json")
|
||||
|
||||
def _meta_path(self) -> str:
|
||||
return os.path.join(self._cache_dir(), _META_FILENAME)
|
||||
|
||||
def _read_meta(self) -> Dict[str, Any]:
|
||||
try:
|
||||
with open(self._meta_path(), "r", encoding="utf-8") as handle:
|
||||
meta = json.load(handle)
|
||||
return meta if isinstance(meta, dict) else {}
|
||||
except FileNotFoundError:
|
||||
return {}
|
||||
except (OSError, json.JSONDecodeError) as exc:
|
||||
logger.warning("Failed to read OpenModelDB cache meta: %s", exc)
|
||||
return {}
|
||||
|
||||
def _read_dump_file(self, path: str) -> Optional[Dict[str, Any]]:
|
||||
try:
|
||||
with open(path, "r", encoding="utf-8") as handle:
|
||||
data = json.load(handle)
|
||||
return data if isinstance(data, dict) else None
|
||||
except (OSError, json.JSONDecodeError) as exc:
|
||||
logger.warning("Failed to read OpenModelDB cache %s: %s", path, exc)
|
||||
return None
|
||||
|
||||
def _write_cache(
|
||||
self,
|
||||
payloads: Dict[str, Dict[str, Any]],
|
||||
etags: Dict[str, str],
|
||||
last_modified: Dict[str, str],
|
||||
) -> None:
|
||||
for name, payload in payloads.items():
|
||||
path = self._dump_path(name)
|
||||
try:
|
||||
with open(path, "w", encoding="utf-8") as handle:
|
||||
json.dump(payload, handle)
|
||||
except OSError as exc:
|
||||
logger.warning("Failed to write OpenModelDB cache %s: %s", path, exc)
|
||||
|
||||
meta = {
|
||||
"fetched_at": time.time(),
|
||||
"etags": etags,
|
||||
"last_modified": last_modified,
|
||||
}
|
||||
try:
|
||||
with open(self._meta_path(), "w", encoding="utf-8") as handle:
|
||||
json.dump(meta, handle, indent=2)
|
||||
except OSError as exc:
|
||||
logger.warning("Failed to write OpenModelDB cache meta: %s", exc)
|
||||
@@ -251,7 +251,28 @@ class ServiceRegistry:
|
||||
cls._services[service_name] = client
|
||||
logger.debug(f"Created and registered {service_name}")
|
||||
return client
|
||||
|
||||
|
||||
@classmethod
|
||||
async def get_openmodeldb_client(cls):
|
||||
"""Get or create OpenModelDB client instance"""
|
||||
service_name = "openmodeldb_client"
|
||||
|
||||
if service_name in cls._services:
|
||||
return cls._services[service_name]
|
||||
|
||||
async with cls._get_lock(service_name):
|
||||
# Double-check after acquiring lock
|
||||
if service_name in cls._services:
|
||||
return cls._services[service_name]
|
||||
|
||||
# Import here to avoid circular imports
|
||||
from .openmodeldb_client import OpenModelDBClient
|
||||
|
||||
client = await OpenModelDBClient.get_instance()
|
||||
cls._services[service_name] = client
|
||||
logger.debug(f"Created and registered {service_name}")
|
||||
return client
|
||||
|
||||
@classmethod
|
||||
async def get_download_manager(cls):
|
||||
"""Get or create Download manager instance"""
|
||||
|
||||
@@ -79,6 +79,9 @@ DEFAULT_SETTINGS: Dict[str, Any] = {
|
||||
"dismissed_banners": [],
|
||||
"enable_metadata_archive_db": False,
|
||||
"enable_civarchive_api": True,
|
||||
# OpenModelDB supplies read-only metadata for upscaler models (the "other"
|
||||
# page's upscaler sub_type) via hash matching against its bulk catalogue.
|
||||
"enable_openmodeldb_api": True,
|
||||
"metadata_provider_order": "civitai_archive_sqlite",
|
||||
"rate_limit_gate_enabled": True,
|
||||
"rate_limit_max_wait_seconds": 300,
|
||||
|
||||
Reference in New Issue
Block a user