From 35f1ced41a6a4422ba7d7aa2a14e777291e6f483 Mon Sep 17 00:00:00 2001 From: Will Miao Date: Sat, 3 Oct 2026 21:13:07 +0800 Subject: [PATCH] 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. --- docs/agent_skills.md | 3 + docs/metadata-json-schema.md | 20 + locales/de.json | 2 + locales/en.json | 2 + locales/es.json | 2 + locales/fr.json | 2 + locales/he.json | 2 + locales/ja.json | 2 + locales/ko.json | 2 + locales/ru.json | 2 + locales/zh-CN.json | 2 + locales/zh-TW.json | 2 + py/routes/handlers/misc_handlers.py | 1 + py/routes/handlers/model_source_handlers.py | 27 +- py/services/metadata_service.py | 32 +- py/services/metadata_sync_service.py | 98 +- py/services/model_metadata_provider.py | 61 +- py/services/model_sources/__init__.py | 4 +- py/services/model_sources/base.py | 38 + py/services/model_sources/openmodeldb.py | 210 +++++ py/services/model_sources/registry.py | 2 + py/services/model_sources/tensorart.py | 1 + py/services/openmodeldb_client.py | 762 ++++++++++++++++ py/services/service_registry.py | 23 +- py/services/settings_manager.py | 3 + static/js/components/shared/ModelModal.js | 16 +- static/js/managers/DownloadManager.js | 4 +- static/js/managers/SettingsManager.js | 5 + static/js/state/index.js | 1 + static/js/utils/modelSourceHelpers.js | 34 +- .../components/modals/settings/library.html | 3 + .../components/modelModal.sourceLinks.test.js | 37 + .../frontend/utils/modelSourceHelpers.test.js | 36 + .../utils/modelSourceUrlDetection.test.js | 11 + tests/routes/test_model_source_handlers.py | 123 +++ tests/services/test_metadata_service.py | 51 ++ tests/services/test_model_sources.py | 4 +- tests/services/test_openmodeldb_client.py | 862 ++++++++++++++++++ tests/services/test_openmodeldb_source.py | 332 +++++++ 39 files changed, 2787 insertions(+), 37 deletions(-) create mode 100644 py/services/model_sources/openmodeldb.py create mode 100644 py/services/openmodeldb_client.py create mode 100644 tests/services/test_openmodeldb_client.py create mode 100644 tests/services/test_openmodeldb_source.py diff --git a/docs/agent_skills.md b/docs/agent_skills.md index 3314dfe4..e9f641a4 100644 --- a/docs/agent_skills.md +++ b/docs/agent_skills.md @@ -74,6 +74,7 @@ Enriches models linked to an external model site with metadata extracted by an L | ModelScope (`modelscope.cn`) | yes | yes | yes | | ModelScope International (`modelscope.ai`) | yes | yes | yes | | TensorArt | yes | no (see below) | no | +| OpenModelDB | yes | yes (card data from the catalogue, no README) | yes | `modelscope.cn` and `modelscope.ai` are **separate catalogues, not mirrors** — a repository published on one is routinely absent from the other — so each is @@ -85,6 +86,8 @@ tables in `modelSourceHelpers.js` and `registry.py` in step when adding a site. TensorArt is link-only: `tensor.art` sits behind a Cloudflare managed challenge and its internal API requires session authorization, so the backend cannot read its model pages. Linking still stores the canonical page URL and the "View on TensorArt" link works. +OpenModelDB is the upscaler catalogue: model ids are flat tokens (no `owner/name`), there are no revisions and no README — `fetch_model_card_context()` reads everything (description, license, tags, example images) from the disk-cached bulk catalogue in `py/services/openmodeldb_client.py`. Only PyTorch resources (`.pth`/`.safetensors`) are downloadable, and only via mirrors that serve raw bytes: HTML-gateway hosts (`mediafire.com`, `mega.nz`, `drive.google.com`) are skipped in favour of a direct mirror, and a model with only gateway mirrors reports a manual-download hint instead of a file list. Filenames are derived from URL path segments (mediafire buries them mid-path) or synthesized as `{model_id}.{type}` for folder links. + **What it does**: 1. Reads the model's `.metadata.json` to get the source (`source_platform` + `source_url`, or the legacy `hf_url`) 2. Fetches the model card through the provider in `py/services/model_sources/` — the README via `fetch_model_card()`, plus any extras the site keeps outside it via `fetch_model_card_context()` diff --git a/docs/metadata-json-schema.md b/docs/metadata-json-schema.md index 39d3e270..066398d5 100644 --- a/docs/metadata-json-schema.md +++ b/docs/metadata-json-schema.md @@ -299,9 +299,29 @@ The `metadata_source` field indicates which provider last updated the metadata: |-------|--------| | `"civitai_api"` | Civitai API | | `"civarchive"` | CivArchive API | +| `"openmodeldb"` | OpenModelDB catalogue (upscaler models only; hash-matched) | | `"archive_db"` | Metadata Archive Database | | `null` | No external source (user-defined only) | +When `metadata_source` is `"openmodeldb"`, the `civitai` payload is a +CivitAI-shaped version dict synthesized from the OpenModelDB catalogue entry +(no numeric `id`/`modelId`), and the OpenModelDB-native details live in its +`openmodeldb` block (`id`, `url`, `authors`, `architecture`, +`architectureName`, `scale`, `license`, `date`). + +In that payload, `images[].url` is always a displayable asset: paired +comparisons use the site-hosted thumbnail because the `LR`/`SR` originals +are ephemeral imgdiff viewer sessions that 404 outside them (the original +viewer link is kept as `images[].meta.comparisonUrl` for reference), and the +model-level thumbnail leads the list since the card preview derives from +`images[0]`. `files[].name` is derived from any URL path segment with a +model extension (mediafire-style mirrors bury it mid-path) or synthesized as +`{model_id}.{type}` for folder links. + +Models downloaded from OpenModelDB additionally carry `source_platform: +"openmodeldb"` and `source_url` (the model page URL); their download-time +hydration is recorded as `metadata_source: "source:openmodeldb"`. + --- ## Auto-Update Behavior diff --git a/locales/de.json b/locales/de.json index 99b92321..17867862 100644 --- a/locales/de.json +++ b/locales/de.json @@ -793,6 +793,8 @@ "downloadComplete": "Download erfolgreich abgeschlossen", "enableCivarchiveApi": "CivArchive API als Metadaten-Anbieter aktivieren", "enableCivarchiveApiHelp": "Wenn aktiviert, wird die CivArchive API als alternative Quelle für Modell-Metadaten verwendet (z. B. für von CivitAI gelöschte Modelle). Deaktivieren, um die Ratenbegrenzungen von CivArchive vollständig zu vermeiden.", + "enableOpenmodeldbApi": "[TODO: Translate] Enable OpenModelDB as metadata provider", + "enableOpenmodeldbApiHelp": "[TODO: Translate] When on, upscaler metadata is also looked up in the OpenModelDB catalogue (openmodeldb.info) when CivitAI has no record. The catalogue is cached locally and refreshed daily.", "providerOrder": "Reihenfolge der Metadaten-Anbieter", "providerOrderHelp": "Die CivitAI API wird immer zuerst versucht. Wählen Sie die Reihenfolge der übrigen Anbieter bei der Metadatensuche.", "providerOrderCivitaiArchiveSqlite": "CivitAI → CivArchive → Archive DB", diff --git a/locales/en.json b/locales/en.json index 9b83bf13..058e83dd 100644 --- a/locales/en.json +++ b/locales/en.json @@ -793,6 +793,8 @@ "downloadComplete": "Download completed successfully", "enableCivarchiveApi": "Enable CivArchive API as metadata provider", "enableCivarchiveApiHelp": "When on, CivArchive API is used as a fallback source for model metadata (e.g. for models deleted from CivitAI). Turn off to avoid CivArchive rate limits entirely.", + "enableOpenmodeldbApi": "Enable OpenModelDB as metadata provider", + "enableOpenmodeldbApiHelp": "When on, upscaler metadata is also looked up in the OpenModelDB catalogue (openmodeldb.info) when CivitAI has no record. The catalogue is cached locally and refreshed daily.", "providerOrder": "Metadata provider fallback order", "providerOrderHelp": "CivitAI API is always tried first. Choose the order of the remaining providers when looking up metadata.", "providerOrderCivitaiArchiveSqlite": "CivitAI → CivArchive → Archive DB", diff --git a/locales/es.json b/locales/es.json index bfc1e6de..2c2d0571 100644 --- a/locales/es.json +++ b/locales/es.json @@ -793,6 +793,8 @@ "downloadComplete": "Descarga completada exitosamente", "enableCivarchiveApi": "Habilitar CivArchive API como proveedor de metadatos", "enableCivarchiveApiHelp": "Al activarlo, la API de CivArchive se usa como fuente alternativa de metadatos de modelos (p. ej. para modelos eliminados de CivitAI). Desactívelo para evitar por completo los límites de velocidad de CivArchive.", + "enableOpenmodeldbApi": "[TODO: Translate] Enable OpenModelDB as metadata provider", + "enableOpenmodeldbApiHelp": "[TODO: Translate] When on, upscaler metadata is also looked up in the OpenModelDB catalogue (openmodeldb.info) when CivitAI has no record. The catalogue is cached locally and refreshed daily.", "providerOrder": "Orden de proveedores de metadatos de respaldo", "providerOrderHelp": "La API de CivitAI siempre se intenta primero. Elija el orden de los demás proveedores al buscar metadatos.", "providerOrderCivitaiArchiveSqlite": "CivitAI → CivArchive → Archive DB", diff --git a/locales/fr.json b/locales/fr.json index 4dbfb441..2bd7e104 100644 --- a/locales/fr.json +++ b/locales/fr.json @@ -793,6 +793,8 @@ "downloadComplete": "Téléchargement terminé avec succès", "enableCivarchiveApi": "Activer l'API CivArchive comme fournisseur de métadonnées", "enableCivarchiveApiHelp": "Lorsqu'elle est activée, l'API CivArchive est utilisée comme source de secours pour les métadonnées des modèles (par ex. pour les modèles supprimés de CivitAI). Désactivez pour éviter entièrement les limites de débit de CivArchive.", + "enableOpenmodeldbApi": "[TODO: Translate] Enable OpenModelDB as metadata provider", + "enableOpenmodeldbApiHelp": "[TODO: Translate] When on, upscaler metadata is also looked up in the OpenModelDB catalogue (openmodeldb.info) when CivitAI has no record. The catalogue is cached locally and refreshed daily.", "providerOrder": "Ordre de secours des fournisseurs de métadonnées", "providerOrderHelp": "L'API CivitAI est toujours essayée en premier. Choisissez l'ordre des autres fournisseurs lors de la recherche de métadonnées.", "providerOrderCivitaiArchiveSqlite": "CivitAI → CivArchive → Archive DB", diff --git a/locales/he.json b/locales/he.json index 7ba591c0..f7862ef5 100644 --- a/locales/he.json +++ b/locales/he.json @@ -793,6 +793,8 @@ "downloadComplete": "ההורדה הושלמה בהצלחה", "enableCivarchiveApi": "הפעל את CivArchive API כספק מטא-נתונים", "enableCivarchiveApiHelp": "כאשר מופעל, CivArchive API משמש כמקור גיבוי למטא-נתונים של מודלים (למשל עבור מודלים שנמחקו מ-CivitAI). כבה כדי להימנע לחלוטין ממגבלות הקצב של CivArchive.", + "enableOpenmodeldbApi": "[TODO: Translate] Enable OpenModelDB as metadata provider", + "enableOpenmodeldbApiHelp": "[TODO: Translate] When on, upscaler metadata is also looked up in the OpenModelDB catalogue (openmodeldb.info) when CivitAI has no record. The catalogue is cached locally and refreshed daily.", "providerOrder": "סדר ספקי מטא-נתונים לגיבוי", "providerOrderHelp": "CivitAI API תמיד מנוסה ראשון. בחר את סדר הספקים הנותרים בעת חיפוש מטא-נתונים.", "providerOrderCivitaiArchiveSqlite": "CivitAI → CivArchive → Archive DB", diff --git a/locales/ja.json b/locales/ja.json index 521ae066..e9596b42 100644 --- a/locales/ja.json +++ b/locales/ja.json @@ -793,6 +793,8 @@ "downloadComplete": "ダウンロードが正常に完了しました", "enableCivarchiveApi": "CivArchive API をメタデータプロバイダーとして有効化", "enableCivarchiveApiHelp": "有効にすると、CivArchive API がモデルメタデータの代替ソースとして使用されます(例:CivitAI から削除されたモデルの場合)。オフにすると、CivArchive のレート制限を完全に回避できます。", + "enableOpenmodeldbApi": "[TODO: Translate] Enable OpenModelDB as metadata provider", + "enableOpenmodeldbApiHelp": "[TODO: Translate] When on, upscaler metadata is also looked up in the OpenModelDB catalogue (openmodeldb.info) when CivitAI has no record. The catalogue is cached locally and refreshed daily.", "providerOrder": "メタデータプロバイダーのフォールバック順序", "providerOrderHelp": "CivitAI API が常に最初に試行されます。メタデータ検索時の残りのプロバイダーの順序を選択してください。", "providerOrderCivitaiArchiveSqlite": "CivitAI → CivArchive → Archive DB", diff --git a/locales/ko.json b/locales/ko.json index fb197ae7..8d393dcb 100644 --- a/locales/ko.json +++ b/locales/ko.json @@ -793,6 +793,8 @@ "downloadComplete": "다운로드가 성공적으로 완료되었습니다", "enableCivarchiveApi": "CivArchive API를 메타데이터 제공자로 활성화", "enableCivarchiveApiHelp": "활성화하면 CivArchive API가 모델 메타데이터의 대체 소스로 사용됩니다 (예: CivitAI에서 삭제된 모델의 경우). 비활성화하면 CivArchive의 속도 제한을 완전히 피할 수 있습니다.", + "enableOpenmodeldbApi": "[TODO: Translate] Enable OpenModelDB as metadata provider", + "enableOpenmodeldbApiHelp": "[TODO: Translate] When on, upscaler metadata is also looked up in the OpenModelDB catalogue (openmodeldb.info) when CivitAI has no record. The catalogue is cached locally and refreshed daily.", "providerOrder": "메타데이터 제공자 폴백 순서", "providerOrderHelp": "CivitAI API가 항상 먼저 시도됩니다. 메타데이터 조회 시 나머지 제공자의 순서를 선택하세요.", "providerOrderCivitaiArchiveSqlite": "CivitAI → CivArchive → Archive DB", diff --git a/locales/ru.json b/locales/ru.json index 353085ae..e6585094 100644 --- a/locales/ru.json +++ b/locales/ru.json @@ -793,6 +793,8 @@ "downloadComplete": "Загрузка успешно завершена", "enableCivarchiveApi": "Включить CivArchive API как источник метаданных", "enableCivarchiveApiHelp": "При включении CivArchive API используется как резервный источник метаданных моделей (например, для моделей, удалённых с CivitAI). Отключите, чтобы полностью избежать ограничений скорости CivArchive.", + "enableOpenmodeldbApi": "[TODO: Translate] Enable OpenModelDB as metadata provider", + "enableOpenmodeldbApiHelp": "[TODO: Translate] When on, upscaler metadata is also looked up in the OpenModelDB catalogue (openmodeldb.info) when CivitAI has no record. The catalogue is cached locally and refreshed daily.", "providerOrder": "Порядок резервных источников метаданных", "providerOrderHelp": "CivitAI API всегда проверяется первым. Выберите порядок остальных источников при поиске метаданных.", "providerOrderCivitaiArchiveSqlite": "CivitAI → CivArchive → Archive DB", diff --git a/locales/zh-CN.json b/locales/zh-CN.json index 01d30be8..8ac7c1c9 100644 --- a/locales/zh-CN.json +++ b/locales/zh-CN.json @@ -793,6 +793,8 @@ "downloadComplete": "下载成功完成", "enableCivarchiveApi": "启用 CivArchive API 作为元数据提供者", "enableCivarchiveApiHelp": "开启后,CivArchive API 将作为模型元数据的备用来源(例如用于已从 CivitAI 删除的模型)。关闭可完全避免 CivArchive 的速率限制。", + "enableOpenmodeldbApi": "[TODO: Translate] Enable OpenModelDB as metadata provider", + "enableOpenmodeldbApiHelp": "[TODO: Translate] When on, upscaler metadata is also looked up in the OpenModelDB catalogue (openmodeldb.info) when CivitAI has no record. The catalogue is cached locally and refreshed daily.", "providerOrder": "元数据提供者回退顺序", "providerOrderHelp": "CivitAI API 始终优先尝试。选择查找元数据时其余提供者的顺序。", "providerOrderCivitaiArchiveSqlite": "CivitAI → CivArchive → Archive DB", diff --git a/locales/zh-TW.json b/locales/zh-TW.json index 4454e8db..97f9087c 100644 --- a/locales/zh-TW.json +++ b/locales/zh-TW.json @@ -793,6 +793,8 @@ "downloadComplete": "下載成功完成", "enableCivarchiveApi": "啟用 CivArchive API 作為中繼資料提供者", "enableCivarchiveApiHelp": "開啟後,CivArchive API 將作為模型中繼資料的備用來源(例如用於已從 CivitAI 刪除的模型)。關閉可完全避免 CivArchive 的速率限制。", + "enableOpenmodeldbApi": "[TODO: Translate] Enable OpenModelDB as metadata provider", + "enableOpenmodeldbApiHelp": "[TODO: Translate] When on, upscaler metadata is also looked up in the OpenModelDB catalogue (openmodeldb.info) when CivitAI has no record. The catalogue is cached locally and refreshed daily.", "providerOrder": "中繼資料提供者回退順序", "providerOrderHelp": "CivitAI API 始終優先嘗試。選擇查詢中繼資料時其餘提供者的順序。", "providerOrderCivitaiArchiveSqlite": "CivitAI → CivArchive → Archive DB", diff --git a/py/routes/handlers/misc_handlers.py b/py/routes/handlers/misc_handlers.py index 21ccbf4c..0acf06ad 100644 --- a/py/routes/handlers/misc_handlers.py +++ b/py/routes/handlers/misc_handlers.py @@ -1807,6 +1807,7 @@ class SettingsHandler: if key in ( "enable_metadata_archive_db", "enable_civarchive_api", + "enable_openmodeldb_api", "metadata_provider_order", ): await self._metadata_provider_updater() diff --git a/py/routes/handlers/model_source_handlers.py b/py/routes/handlers/model_source_handlers.py index db6fa07f..d227e142 100644 --- a/py/routes/handlers/model_source_handlers.py +++ b/py/routes/handlers/model_source_handlers.py @@ -31,7 +31,6 @@ from ...services.model_sources import ( detect_source, get_download_source, hydrate_from_source, - is_valid_source_id, list_sources, normalize_metadata_source, ) @@ -277,9 +276,7 @@ class ModelSourceHandler: "supports_enrichment": source.supports_enrichment, "supports_download": source.supports_download, "default_revision": source.default_revision, - "example_url": source.canonical_url( - "user/repo" if source.platform != "tensorart" else "827823520299086029" - ), + "example_url": source.canonical_url(source.example_source_id), } for source in list_sources() ]) @@ -326,9 +323,7 @@ class ModelSourceHandler: "error": ( "Unsupported model URL. Supported formats: " + ", ".join( - f"{s.label} ({s.canonical_url('user/repo')})" - if s.platform != "tensorart" - else f"{s.label} (https://tensor.art/models/)" + f"{s.label} ({s.canonical_url(s.example_source_id)})" for s in list_sources() ) ), @@ -424,9 +419,9 @@ class ModelSourceHandler: source = get_download_source(platform) if source is None: return _unsupported_platform_error(platform) - if not is_valid_source_id(repo): + if not source.is_valid_source_id(repo): return web.json_response( - {"error": "Missing or invalid 'repo' parameter (expected owner/name)"}, + {"error": "Missing or invalid 'repo' parameter"}, status=400, ) @@ -492,10 +487,11 @@ class ModelSourceHandler: {"error": "Missing required fields: 'repo' and 'filename'"}, status=400 ) - # `owner/name` only; the components become path segments below. - if not is_valid_source_id(repo): + # The id becomes a path segment below; each site defines what a safe + # id looks like (`owner/name` for repository sites, a flat token for + # OpenModelDB). + if not source.is_valid_source_id(repo): return web.json_response({"error": f"Invalid repo format: {repo}"}, status=400) - owner, repo_name = repo.split("/", 1) # Validate filename — must not contain path traversal if ".." in filename: @@ -521,7 +517,7 @@ class ModelSourceHandler: base_dir = os.path.normpath(os.path.join(os.getcwd(), "models", model_root)) if use_default_paths: - target_dir = os.path.join(base_dir, source.default_subdir, owner, repo_name) + target_dir = os.path.join(base_dir, *source.default_subdir_parts(repo)) elif relative_path: target_dir = os.path.join(base_dir, relative_path) else: @@ -536,7 +532,10 @@ class ModelSourceHandler: # Built per request: sites that redirect to a CDN hand out a # time-limited token in the redirect, so the URL must never be cached. - resolve_url = source.file_download_url(repo, filename, revision) + try: + resolve_url = await source.resolve_download_url(repo, filename, revision) + except ModelSourceError as exc: + return web.json_response({"error": str(exc)}, status=exc.status) ref = SourceRef( platform=source.platform, source_id=repo, url=source.canonical_url(repo) ) diff --git a/py/services/metadata_service.py b/py/services/metadata_service.py index 9024a77e..e38e2470 100644 --- a/py/services/metadata_service.py +++ b/py/services/metadata_service.py @@ -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 diff --git a/py/services/metadata_sync_service.py b/py/services/metadata_sync_service.py index 337ec381..c69c47c7 100644 --- a/py/services/metadata_sync_service.py +++ b/py/services/metadata_sync_service.py @@ -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) diff --git a/py/services/model_metadata_provider.py b/py/services/model_metadata_provider.py index a45579bf..9676c6a3 100644 --- a/py/services/model_metadata_provider.py +++ b/py/services/model_metadata_provider.py @@ -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) diff --git a/py/services/model_sources/__init__.py b/py/services/model_sources/__init__.py index 9a412754..8cfae4c4 100644 --- a/py/services/model_sources/__init__.py +++ b/py/services/model_sources/__init__.py @@ -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", diff --git a/py/services/model_sources/base.py b/py/services/model_sources/base.py index afd601a7..aeb6ac80 100644 --- a/py/services/model_sources/base.py +++ b/py/services/model_sources/base.py @@ -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 ``//`` 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.""" diff --git a/py/services/model_sources/openmodeldb.py b/py/services/model_sources/openmodeldb.py new file mode 100644 index 00000000..d117d920 --- /dev/null +++ b/py/services/model_sources/openmodeldb.py @@ -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[A-Za-z0-9_][A-Za-z0-9_.\-]*)" +) +_STRICT_URL_PATTERN = re.compile( + r"https?://(?:www\.)?openmodeldb\.info/models/(?P[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"] diff --git a/py/services/model_sources/registry.py b/py/services/model_sources/registry.py index df568c42..f2b338b1 100644 --- a/py/services/model_sources/registry.py +++ b/py/services/model_sources/registry.py @@ -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} diff --git a/py/services/model_sources/tensorart.py b/py/services/model_sources/tensorart.py index 8d68a254..d9bcd0dc 100644 --- a/py/services/model_sources/tensorart.py +++ b/py/services/model_sources/tensorart.py @@ -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 diff --git a/py/services/openmodeldb_client.py b/py/services/openmodeldb_client.py new file mode 100644 index 00000000..17bf7e0b --- /dev/null +++ b/py/services/openmodeldb_client.py @@ -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 ```` 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//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) diff --git a/py/services/service_registry.py b/py/services/service_registry.py index 5c0a7c37..99b9e997 100644 --- a/py/services/service_registry.py +++ b/py/services/service_registry.py @@ -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""" diff --git a/py/services/settings_manager.py b/py/services/settings_manager.py index df564d30..594844ab 100644 --- a/py/services/settings_manager.py +++ b/py/services/settings_manager.py @@ -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, diff --git a/static/js/components/shared/ModelModal.js b/static/js/components/shared/ModelModal.js index 73a84515..1cec9af8 100644 --- a/static/js/components/shared/ModelModal.js +++ b/static/js/components/shared/ModelModal.js @@ -1,5 +1,5 @@ import { showToast, openCivitai, sendLoraToWorkflow, sendEmbeddingToWorkflow, sendModelPathToWorkflow, buildLoraSyntax, copyToClipboard } from '../../utils/uiHelpers.js'; -import { getModelSourceInfo, getModelSourceGroupKey, getModelSourceViewTitle, openModelSource } from '../../utils/modelSourceHelpers.js'; +import { getModelSource, getModelSourceInfo, getModelSourceGroupKey, getModelSourceViewTitle, openModelSource } from '../../utils/modelSourceHelpers.js'; import { modalManager } from '../../managers/ModalManager.js'; import { MODEL_TYPES } from '../../api/apiConfig.js'; import { @@ -32,6 +32,18 @@ function getModalFilePath(fallback = '') { return fallback; } +/** + * Source descriptor for a model that was hash-enriched from the OpenModelDB + * catalogue: it carries no `source_url`, so `getModelSourceInfo` finds + * nothing and the page link lives in the civitai payload instead. + */ +function getOpenModelDBSourceInfo(model) { + const url = model?.civitai?.openmodeldb?.url; + if (typeof url !== 'string' || !url) return null; + const descriptor = getModelSource('openmodeldb'); + return descriptor ? { ...descriptor, sourceId: '', url } : null; +} + const COMMERCIAL_ICON_CONFIG = [ { key: 'image', @@ -398,7 +410,7 @@ export async function showModelModal(model, modelType) {
${translate('modals.model.actions.viewOnCivitaiText', {}, 'View on Civitai')}
`.trim() : ''; - const sourceInfo = getModelSourceInfo(modelWithFullData); + const sourceInfo = getModelSourceInfo(modelWithFullData) || getOpenModelDBSourceInfo(modelWithFullData); const escapedSourceUrl = sourceInfo?.url ? escapeAttribute(sourceInfo.url) : ''; const isHuggingFaceSource = sourceInfo?.platform === 'huggingface'; const sourceTitle = sourceInfo ? getModelSourceViewTitle(sourceInfo) : ''; diff --git a/static/js/managers/DownloadManager.js b/static/js/managers/DownloadManager.js index 7520a958..08c0b31d 100644 --- a/static/js/managers/DownloadManager.js +++ b/static/js/managers/DownloadManager.js @@ -18,7 +18,7 @@ import { detectModelSourceDownloadUrl, getModelSource, isExternalModelSource, - isValidRepoId, + isValidSourceId, } from '../utils/modelSourceHelpers.js'; export class DownloadManager { @@ -534,7 +534,7 @@ export class DownloadManager { const sourceInfo = detectModelSourceDownloadUrl(trimmed); if (sourceInfo) { // Reject path-traversal patterns like "../.." or "user/.." - if (!isValidRepoId(sourceInfo.repo)) { + if (!isValidSourceId(sourceInfo.platform, sourceInfo.repo)) { return null; } return { diff --git a/static/js/managers/SettingsManager.js b/static/js/managers/SettingsManager.js index e6763d97..2f9486f0 100644 --- a/static/js/managers/SettingsManager.js +++ b/static/js/managers/SettingsManager.js @@ -3711,6 +3711,11 @@ export class SettingsManager { enableCivarchiveApiCheckbox.checked = state.global.settings.enable_civarchive_api ?? true; } + const enableOpenmodeldbApiCheckbox = document.getElementById('enableOpenmodeldbApi'); + if (enableOpenmodeldbApiCheckbox) { + enableOpenmodeldbApiCheckbox.checked = state.global.settings.enable_openmodeldb_api ?? true; + } + const metadataProviderOrderSelect = document.getElementById('metadataProviderOrder'); if (metadataProviderOrderSelect) { metadataProviderOrderSelect.value = state.global.settings.metadata_provider_order || 'civitai_archive_sqlite'; diff --git a/static/js/state/index.js b/static/js/state/index.js index acf76dd3..6d04dcd3 100644 --- a/static/js/state/index.js +++ b/static/js/state/index.js @@ -16,6 +16,7 @@ const DEFAULT_SETTINGS_BASE = Object.freeze({ show_only_sfw: false, enable_metadata_archive_db: false, enable_civarchive_api: true, + enable_openmodeldb_api: true, metadata_provider_order: 'civitai_archive_sqlite', proxy_enabled: false, proxy_type: 'http', diff --git a/static/js/utils/modelSourceHelpers.js b/static/js/utils/modelSourceHelpers.js index c23e2c0c..39e43b26 100644 --- a/static/js/utils/modelSourceHelpers.js +++ b/static/js/utils/modelSourceHelpers.js @@ -1,5 +1,5 @@ /** - * External model source helpers (Hugging Face / ModelScope / TensorArt). + * External model source helpers (Hugging Face / ModelScope / TensorArt / OpenModelDB). * * Mirrors `py/services/model_sources/registry.py` so the frontend and the * backend agree on URL recognition, version-group keys, and which sites @@ -92,6 +92,26 @@ export const MODEL_SOURCES = [ canonical: (id) => `https://tensor.art/models/${id}`, filePage: null, }, + { + // Flat catalogue of upscalers: the id IS the published-model identity + // (no owner/repo split, no revisions, no per-file page). + platform: 'openmodeldb', + label: 'OpenModelDB', + groupPrefix: 'omdb', + groupKey: 'repo', + supportsEnrichment: true, + supportsDownload: true, + defaultRevision: '', + defaultSubdir: 'openmodeldb', + exampleUrl: 'https://openmodeldb.info/models/4x-UltraSharp', + placeholder: 'https://openmodeldb.info/models/4x-UltraSharp', + pattern: /^https?:\/\/(?:www\.)?openmodeldb\.info\/models\/([A-Za-z0-9_][A-Za-z0-9_.-]*)/i, + filePattern: null, + // Flat model ids (no owner/name split). + flatId: true, + canonical: (id) => `https://openmodeldb.info/models/${id}`, + filePage: null, + }, ]; /** Return the source descriptor for a platform id, or null. */ @@ -257,6 +277,18 @@ export function isValidRepoId(repo) { .every((part) => part && part !== '.' && part !== '..' && /^[A-Za-z0-9_][\w.-]*$/.test(part)); } +/** + * Source-aware id validation, mirroring `ModelSource.is_valid_source_id`: + * `owner/name` for repository sites, a flat token for OpenModelDB. + */ +export function isValidSourceId(platform, repo) { + const source = getModelSource(platform); + if (source?.flatId) { + return typeof repo === 'string' && /^[A-Za-z0-9_][A-Za-z0-9_.-]*$/.test(repo); + } + return isValidRepoId(repo); +} + /** * Recognise a downloadable model-source URL. * diff --git a/templates/components/modals/settings/library.html b/templates/components/modals/settings/library.html index 2b50fa3d..3a342d6b 100644 --- a/templates/components/modals/settings/library.html +++ b/templates/components/modals/settings/library.html @@ -292,6 +292,9 @@ {{ sm.setting_toggle('enableCivarchiveApi', 'enable_civarchive_api', 'settings.metadataArchive.enableCivarchiveApi', 'settings.metadataArchive.enableCivarchiveApiHelp') }} + + {{ sm.setting_toggle('enableOpenmodeldbApi', 'enable_openmodeldb_api', 'settings.metadataArchive.enableOpenmodeldbApi', 'settings.metadataArchive.enableOpenmodeldbApiHelp') }} + {{ sm.setting_toggle('enableMetadataArchive', 'enable_metadata_archive_db', 'settings.metadataArchive.enableArchiveDb', 'settings.metadataArchive.enableArchiveDbHelp') }} diff --git a/tests/frontend/components/modelModal.sourceLinks.test.js b/tests/frontend/components/modelModal.sourceLinks.test.js index 9eef1dc0..354c9041 100644 --- a/tests/frontend/components/modelModal.sourceLinks.test.js +++ b/tests/frontend/components/modelModal.sourceLinks.test.js @@ -188,4 +188,41 @@ describe('Model modal source links (#1094)', () => { expect(civitaiLink()).toBeNull(); expect(hfLink()).toBeNull(); }); + + it('renders an OpenModelDB link from the enriched civitai payload', async () => { + // Hash-enriched upscalers carry no `source_url`; the page link lives in + // `civitai.openmodeldb.url` instead. + await renderModal( + makeModel({ + from_civitai: false, + civitai: { + source: 'openmodeldb', + openmodeldb: { id: '4x-UltraSharp', url: 'https://openmodeldb.info/models/4x-UltraSharp' }, + }, + }) + ); + + const link = document.querySelector('[data-action="view-model-source"]'); + expect(link).not.toBeNull(); + expect(link.dataset.sourceUrl).toBe('https://openmodeldb.info/models/4x-UltraSharp'); + expect(civitaiLink()).toBeNull(); + }); + + it('prefers source_url over the openmodeldb payload fallback', async () => { + await renderModal( + makeModel({ + from_civitai: false, + source_url: 'https://openmodeldb.info/models/1x-DeJPG', + source_platform: 'openmodeldb', + civitai: { + source: 'openmodeldb', + openmodeldb: { id: '4x-UltraSharp', url: 'https://openmodeldb.info/models/4x-UltraSharp' }, + }, + }) + ); + + const links = document.querySelectorAll('[data-action="view-model-source"]'); + expect(links.length).toBe(1); + expect(links[0].dataset.sourceUrl).toBe('https://openmodeldb.info/models/1x-DeJPG'); + }); }); diff --git a/tests/frontend/utils/modelSourceHelpers.test.js b/tests/frontend/utils/modelSourceHelpers.test.js index 629aa896..2f7ad0aa 100644 --- a/tests/frontend/utils/modelSourceHelpers.test.js +++ b/tests/frontend/utils/modelSourceHelpers.test.js @@ -18,6 +18,7 @@ const { canEnrichModelSource, getModelSourceViewTitle, parseModelSourceGroupKey, + isValidSourceId, openModelSource, } = await import('../../../static/js/utils/modelSourceHelpers.js'); @@ -28,6 +29,7 @@ describe('modelSourceHelpers', () => { 'modelscope', 'modelscope-ai', 'tensorart', + 'openmodeldb', ]); }); @@ -64,6 +66,16 @@ describe('modelSourceHelpers', () => { expect(info.url).toBe('https://tensor.art/models/827823520299086029'); }); + it('recognises OpenModelDB URLs with flat model ids', () => { + const info = parseModelSourceUrl('https://openmodeldb.info/models/4x-UltraSharp'); + expect(info.platform).toBe('openmodeldb'); + expect(info.groupPrefix).toBe('omdb'); + expect(info.sourceId).toBe('4x-UltraSharp'); + expect(info.url).toBe('https://openmodeldb.info/models/4x-UltraSharp'); + expect(info.supportsDownload).toBe(true); + expect(info.supportsEnrichment).toBe(true); + }); + it('rejects unsupported URLs', () => { expect(parseModelSourceUrl('https://example.com/x')).toBeNull(); expect(parseModelSourceUrl('')).toBeNull(); @@ -109,6 +121,10 @@ describe('modelSourceHelpers', () => { expect(getModelSourceGroupKey({ source_url: 'https://tensor.art/models/123' })).toBe( 'ta:123' ); + // OpenModelDB's flat model id is the published-model identity. + expect( + getModelSourceGroupKey({ source_url: 'https://openmodeldb.info/models/4x-UltraSharp' }) + ).toBe('omdb:4x-UltraSharp'); // ModelScope groups by the site-native published-model id. expect( getModelSourceGroupKey({ @@ -181,6 +197,11 @@ describe('modelSourceHelpers', () => { }); expect(parseModelSourceGroupKey('ms:user/repo').platform).toBe('modelscope'); expect(parseModelSourceGroupKey('ta:123').platform).toBe('tensorart'); + expect(parseModelSourceGroupKey('omdb:4x-UltraSharp')).toEqual({ + platform: 'openmodeldb', + label: 'OpenModelDB', + sourceId: '4x-UltraSharp', + }); }); it('rejects numeric CivitAI model ids and unknown prefixes', () => { @@ -192,6 +213,21 @@ describe('modelSourceHelpers', () => { }); }); + describe('isValidSourceId', () => { + it('requires owner/name for repository sites', () => { + expect(isValidSourceId('huggingface', 'user/repo')).toBe(true); + expect(isValidSourceId('huggingface', '4x-UltraSharp')).toBe(false); + expect(isValidSourceId('modelscope', 'u/..')).toBe(false); + }); + + it('accepts flat model ids for OpenModelDB', () => { + expect(isValidSourceId('openmodeldb', '4x-UltraSharp')).toBe(true); + expect(isValidSourceId('openmodeldb', 'owner/name')).toBe(false); + expect(isValidSourceId('openmodeldb', '../escape')).toBe(false); + expect(isValidSourceId('openmodeldb', '')).toBe(false); + }); + }); + describe('openModelSource', () => { it('opens the URL in a new tab', () => { const openSpy = vi.spyOn(window, 'open').mockImplementation(() => {}); diff --git a/tests/frontend/utils/modelSourceUrlDetection.test.js b/tests/frontend/utils/modelSourceUrlDetection.test.js index ff35b311..81162d66 100644 --- a/tests/frontend/utils/modelSourceUrlDetection.test.js +++ b/tests/frontend/utils/modelSourceUrlDetection.test.js @@ -222,6 +222,17 @@ describe('DownloadManager.detectUrlType — external model source URLs', () => { expect(intl.platform).toBe('modelscope-ai'); }); + it('detects an OpenModelDB model URL with its flat id', () => { + const result = DownloadManager.detectUrlType( + 'https://openmodeldb.info/models/4x-UltraSharp' + ); + expect(result).toEqual({ + type: 'model-source-repo', + platform: 'openmodeldb', + repo: '4x-UltraSharp', + }); + }); + it('rejects path traversal in either platform', () => { expect( DownloadManager.detectUrlType('https://modelscope.cn/models/../etc/passwd') diff --git a/tests/routes/test_model_source_handlers.py b/tests/routes/test_model_source_handlers.py index b3e3ed80..12c58d61 100644 --- a/tests/routes/test_model_source_handlers.py +++ b/tests/routes/test_model_source_handlers.py @@ -312,6 +312,7 @@ async def test_get_model_sources_lists_capabilities(): "modelscope", "modelscope-ai", "tensorart", + "openmodeldb", } assert by_platform["huggingface"]["supports_enrichment"] is True assert by_platform["modelscope"]["supports_enrichment"] is True @@ -1305,3 +1306,125 @@ async def test_download_model_source_sends_no_headers_without_hf_token( assert response.status == 200 assert captured["custom_headers"] is None + + +# --------------------------------------------------------------------------- +# OpenModelDB downloads +# --------------------------------------------------------------------------- + + +def _seed_openmodeldb_client(tmp_path, monkeypatch) -> None: + """Install a catalogue-loaded OpenModelDB client as the singleton.""" + from py.services.openmodeldb_client import OpenModelDBClient + + client = OpenModelDBClient(cache_dir=str(tmp_path / "omdb-cache")) + client._install_payloads( + { + "models": { + "4x-UltraSharp": { + "name": "4x UltraSharp", + "resources": [ + { + "platform": "pytorch", + "type": "pth", + "size": 67_000_000, + "sha256": "a" * 64, + "urls": ["https://files.example.com/4x-UltraSharp.pth"], + } + ], + } + }, + "users": {}, + "tags": {}, + "architectures": {}, + } + ) + monkeypatch.setattr( + OpenModelDBClient, "get_instance", AsyncMock(return_value=client) + ) + + +@pytest.mark.asyncio +async def test_download_model_source_openmodeldb_default_paths(tmp_path, monkeypatch): + captured = _stub_download_backend(monkeypatch) + saved = AsyncMock() + monkeypatch.setattr(model_source_handlers, "_save_source_metadata", saved) + _seed_openmodeldb_client(tmp_path, monkeypatch) + + response = await ModelSourceHandler().download_model_source( + FakeRequest( + json_data={ + "platform": "openmodeldb", + # Flat catalogue id — no owner/name split. + "repo": "4x-UltraSharp", + "filename": "4x-UltraSharp.pth", + "model_root": str(tmp_path), + "use_default_paths": True, + } + ) + ) + + assert response.status == 200 + assert captured["url"] == "https://files.example.com/4x-UltraSharp.pth" + # Flat layout: the site sub-directory only, no owner/repo namespaces. + assert captured["save_path"] == str( + tmp_path / "openmodeldb" / "4x-UltraSharp.pth" + ) + + ref = saved.await_args.args[1] + assert ref.platform == "openmodeldb" + assert ref.source_id == "4x-UltraSharp" + assert ref.url == "https://openmodeldb.info/models/4x-UltraSharp" + + +@pytest.mark.asyncio +async def test_download_model_source_openmodeldb_unknown_model_returns_404( + tmp_path, monkeypatch +): + _stub_download_backend(monkeypatch) + _seed_openmodeldb_client(tmp_path, monkeypatch) + + response = await ModelSourceHandler().download_model_source( + FakeRequest( + json_data={ + "platform": "openmodeldb", + "repo": "nope", + "filename": "f.pth", + "model_root": str(tmp_path), + } + ) + ) + + assert response.status == 404 + assert "not found" in _json_payload(response)["error"] + + +@pytest.mark.asyncio +async def test_download_model_source_openmodeldb_rejects_repo_style_id(tmp_path): + response = await ModelSourceHandler().download_model_source( + FakeRequest( + json_data={ + "platform": "openmodeldb", + "repo": "owner/name", + "filename": "f.pth", + "model_root": str(tmp_path), + } + ) + ) + + assert response.status == 400 + assert "Invalid repo format" in _json_payload(response)["error"] + + +@pytest.mark.asyncio +async def test_list_model_source_files_openmodeldb(tmp_path, monkeypatch): + _seed_openmodeldb_client(tmp_path, monkeypatch) + + response = await ModelSourceHandler().list_model_source_files( + FakeRequest(query={"platform": "openmodeldb", "repo": "4x-UltraSharp"}) + ) + + assert response.status == 200 + assert _json_payload(response) == [ + {"filename": "4x-UltraSharp.pth", "size": 67_000_000} + ] diff --git a/tests/services/test_metadata_service.py b/tests/services/test_metadata_service.py index 95f8bfdf..dc6e0528 100644 --- a/tests/services/test_metadata_service.py +++ b/tests/services/test_metadata_service.py @@ -73,6 +73,7 @@ def _stub_settings(**overrides): base = { "enable_metadata_archive_db": False, "enable_civarchive_api": True, + "enable_openmodeldb_api": False, "metadata_provider_order": "civitai_archive_sqlite", } base.update(overrides) @@ -99,6 +100,11 @@ async def _run_initialize(monkeypatch, settings): "get_civarchive_client", AsyncMock(return_value=object()), ) + monkeypatch.setattr( + metadata_service.ServiceRegistry, + "get_openmodeldb_client", + AsyncMock(return_value=object()), + ) # Make MetadataArchiveManager report a usable db path when enabled fake_archive = SimpleNamespace(get_database_path=lambda: "/tmp/fake.db") @@ -172,3 +178,48 @@ async def test_initialize_providers_single_provider_when_only_civitai(monkeypatc assert "fallback" not in manager.providers assert manager.default_provider == "civitai_api" + + +# --------------------------------------------------------------------------- +# initialize_metadata_providers — OpenModelDB gating + ordering +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_initialize_providers_includes_openmodeldb_when_enabled(monkeypatch): + settings = _stub_settings(enable_openmodeldb_api=True) + manager = await _run_initialize(monkeypatch, settings) + assert "openmodeldb_api" in manager.providers + # Local-index lookups run before the rate-limited CivArchive network API. + assert _fallback_provider_order(manager) == [ + "civitai_api", + "openmodeldb_api", + "civarchive_api", + ] + + +@pytest.mark.asyncio +async def test_initialize_providers_openmodeldb_sqlite_preset(monkeypatch): + settings = _stub_settings( + enable_metadata_archive_db=True, + enable_openmodeldb_api=True, + metadata_provider_order="civitai_sqlite_archive", + ) + manager = await _run_initialize(monkeypatch, settings) + assert _fallback_provider_order(manager) == [ + "civitai_api", + "openmodeldb_api", + "sqlite", + "civarchive_api", + ] + + +@pytest.mark.asyncio +async def test_initialize_providers_skips_openmodeldb_when_disabled(monkeypatch): + settings = _stub_settings( + enable_metadata_archive_db=True, + enable_openmodeldb_api=False, + ) + manager = await _run_initialize(monkeypatch, settings) + assert "openmodeldb_api" not in manager.providers + assert _fallback_provider_order(manager) == ["civitai_api", "civarchive_api", "sqlite"] diff --git a/tests/services/test_model_sources.py b/tests/services/test_model_sources.py index 31819d5a..74a5ee3d 100644 --- a/tests/services/test_model_sources.py +++ b/tests/services/test_model_sources.py @@ -195,6 +195,7 @@ class TestCapabilities: "modelscope", "modelscope-ai", "tensorart", + "openmodeldb", } def test_labels_are_brand_names(self): @@ -202,6 +203,7 @@ class TestCapabilities: assert source_label("modelscope") == "ModelScope" assert source_label("modelscope-ai") == "ModelScope (International)" assert source_label("tensorart") == "TensorArt" + assert source_label("openmodeldb") == "OpenModelDB" assert source_label("unknown", "fallback") == "fallback" @@ -920,7 +922,7 @@ class TestSourceIdValidation: class TestDownloadSourceRegistry: def test_downloadable_sources_excludes_link_only_sites(self): platforms = {source.platform for source in downloadable_sources()} - assert platforms == {"huggingface", "modelscope", "modelscope-ai"} + assert platforms == {"huggingface", "modelscope", "modelscope-ai", "openmodeldb"} def test_get_download_source_rejects_link_only_platform(self): assert get_download_source("tensorart") is None diff --git a/tests/services/test_openmodeldb_client.py b/tests/services/test_openmodeldb_client.py new file mode 100644 index 00000000..b0fcfc8b --- /dev/null +++ b/tests/services/test_openmodeldb_client.py @@ -0,0 +1,862 @@ +"""Tests for the OpenModelDB client, provider, and sync-service integration.""" + +from __future__ import annotations + +import json +import time +from types import SimpleNamespace +from typing import Any, Dict, Optional +from unittest.mock import AsyncMock + +import pytest + +from py.services import openmodeldb_client as openmodeldb_module +from py.services.metadata_sync_service import MetadataSyncService +from py.services.model_metadata_provider import ( + FallbackMetadataProvider, + ModelMetadataProvider, + OpenModelDBModelMetadataProvider, +) +from py.services.openmodeldb_client import ( + CACHE_TTL_SECONDS, + OPENMODELDB_API_BASE, + OpenModelDBClient, +) + + +MODEL_SHA = "a" * 64 +ALT_SHA = "b" * 64 + + +def _fixture_dumps() -> Dict[str, Dict[str, Any]]: + """Small OpenModelDB catalogue fixture covering the shape variants.""" + models = { + "4x-UltraSharp": { + "name": "4x UltraSharp", + "author": "kim", + "license": "MIT", + "tags": ["general"], + "description": "Sharp general-purpose upscaler", + "date": "2023-01-02", + "architecture": "esrgan", + "size": ["64nf", "23nb"], + "scale": 4, + "inputChannels": 3, + "outputChannels": 3, + "thumbnail": { + "type": "paired", + "LR": "/thumbs/ultrasharp-lr.png", + "SR": "/thumbs/ultrasharp-sr.jpg", + }, + "resources": [ + { + "platform": "pytorch", + "type": "pth", + "size": 67_000_000, + "sha256": MODEL_SHA, + "urls": ["https://files.example.com/4x-UltraSharp.pth"], + }, + { + "platform": "pytorch", + "type": "safetensors", + "size": 67_100_000, + "sha256": ALT_SHA, + "urls": ["https://files.example.com/4x-UltraSharp.safetensors"], + }, + ], + "images": [ + { + "type": "paired", + "caption": "comparison", + "LR": "https://img.example.com/lr.jpg", + "SR": "https://img.example.com/sr.jpg", + "thumbnail": "/thumbs/small/t.jpg", + } + ], + }, + "2x-90s-Sonic-LG": { + "name": "90s Sonic 2x (Large)", + "author": "kim", + "architecture": "esrgan", + "scale": 2, + "thumbnail": { + "type": "paired", + "LR": "/thumbs/sonic-lr.png", + "SR": "/thumbs/sonic-sr.jpg", + }, + "resources": [ + { + "platform": "pytorch", + "type": "pth", + "size": 9_443_322, + "sha256": "d" * 64, + # mediafire puts the real filename mid-path; the basename + # is just "file". + "urls": [ + "https://www.mediafire.com/file/e57pd40qph8nak2/90s_Sonic_2x.pth/file" + ], + } + ], + "images": [ + { + "type": "paired", + "caption": "Tails", + # Ephemeral imgdiff viewer sessions: 404 outside them. + "LR": "https://imgdiff.net/api/image.php?id=abc&image=2171", + "SR": "https://imgdiff.net/api/image.php?id=abc&image=2172", + "thumbnail": "/thumbs/small/sonic1.jpg", + }, + { + "type": "paired", + "LR": "https://imgdiff.net/api/image.php?id=def&image=2173", + "SR": "https://imgdiff.net/api/image.php?id=def&image=2174", + }, + ], + }, + "mega-only": { + "name": "Mega Only", + "author": "kim", + "architecture": "esrgan", + "scale": 4, + "resources": [ + { + "platform": "pytorch", + "type": "pth", + "size": 60_000_000, + "sha256": "e" * 64, + # A folder link carries no filename at all. + "urls": ["https://mega.nz/folder/qZRBmaIY#nIG8KyWFcGNTuMX_XNbJ_g"], + } + ], + "images": [], + }, + "1x-DeJPG": { + "name": "1x DeJPG", + "author": ["alice", "bob"], + "license": None, + "tags": ["restoration"], + "description": "", + "date": "2021-05-06", + "architecture": "compact", + "size": [], + "scale": 1, + "inputChannels": 3, + "outputChannels": 3, + "resources": [ + { + "platform": "pytorch", + "type": "pth", + "size": 3_000_000, + "sha256": "c" * 64, + "urls": ["https://files.example.com/1x-DeJPG.pth"], + } + ], + "images": [ + { + "type": "standalone", + "url": "https://img.example.com/dejpg.png", + } + ], + }, + } + users = { + "kim": {"name": "Kim"}, + "alice": {"name": "Alice"}, + "bob": {"name": "Bob"}, + } + tags = { + "general": {"name": "General Purpose"}, + "restoration": {"name": "Restoration"}, + } + architectures = { + "esrgan": {"name": "ESRGAN"}, + "compact": {"name": "Compact"}, + } + return { + "models": models, + "users": users, + "tags": tags, + "architectures": architectures, + } + + +class DummyDownloader: + """Downloader stub exposing the two entry points the client uses.""" + + def __init__( + self, + *, + payloads: Optional[Dict[str, Any]] = None, + headers: Optional[Dict[str, Dict[str, str]]] = None, + fail: bool = False, + ) -> None: + self.payloads = payloads or {} + self.headers = headers or {} + self.fail = fail + self.get_calls: list = [] + self.head_calls: list = [] + + async def get_response_headers(self, url, use_auth=False, custom_headers=None): + self.head_calls.append(url) + if self.fail: + return False, "network unreachable" + return True, dict(self.headers.get(url, {})) + + async def make_request(self, method, url, use_auth=False, custom_headers=None, **kwargs): + self.get_calls.append(url) + if self.fail: + return False, "network unreachable" + return True, self.payloads.get(url) + + +def _dump_url(name: str) -> str: + return f"{OPENMODELDB_API_BASE}/{name}.json" + + +@pytest.fixture +def dumps() -> Dict[str, Dict[str, Any]]: + return _fixture_dumps() + + +@pytest.fixture +def make_client(monkeypatch, tmp_path): + """Factory returning a client backed by a DummyDownloader and tmp cache.""" + + def _make(downloader: DummyDownloader, **kwargs) -> OpenModelDBClient: + monkeypatch.setattr( + openmodeldb_module, + "get_downloader", + AsyncMock(return_value=downloader), + ) + return OpenModelDBClient(cache_dir=str(tmp_path / "omdb"), **kwargs) + + return _make + + +def _payloads_by_url(dumps: Dict[str, Dict[str, Any]]) -> Dict[str, Any]: + return {f"{OPENMODELDB_API_BASE}/{name}.json": payload for name, payload in dumps.items()} + + +# --------------------------------------------------------------------------- +# Index building + CivitAI-shape mapping +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_lookup_maps_catalogue_entry_to_civitai_shape(make_client, dumps): + downloader = DummyDownloader(payloads=_payloads_by_url(dumps)) + client = make_client(downloader) + + result, error = await client.get_model_by_hash(MODEL_SHA) + + assert error is None + assert result is not None + + # Version-level fields + assert result["name"] == "4x UltraSharp" + assert result["source"] == "openmodeldb" + assert result["baseModel"] == "Other" + assert result["publishedAt"] == "2023-01-02" + assert result["trainedWords"] == [] + assert result["description"] == "Sharp general-purpose upscaler" + # OpenModelDB ids are strings; numeric CivitAI ids must stay absent so + # consumers do not treat this as a CivitAI model. + assert "id" not in result + assert "modelId" not in result + + # Model block + model = result["model"] + assert model["name"] == "4x UltraSharp" + assert model["type"] == "Upscaler" + assert model["tags"] == ["General Purpose"] + assert model["license"] == "MIT" + assert model["description"] == "Sharp general-purpose upscaler" + + # Creator resolved via users.json + assert result["creator"]["username"] == "Kim" + + # Files: every resource becomes a file; the hash-matched one is primary. + files = result["files"] + assert len(files) == 2 + primary = [f for f in files if f["primary"]] + assert len(primary) == 1 + assert primary[0]["name"] == "4x-UltraSharp.pth" + assert primary[0]["hashes"] == {"SHA256": MODEL_SHA.upper()} + assert primary[0]["sizeKB"] == pytest.approx(67_000_000 / 1024.0) + assert primary[0]["downloadUrl"] == "https://files.example.com/4x-UltraSharp.pth" + assert primary[0]["metadata"]["format"] == "PickleTensor" + st_file = next(f for f in files if f["name"].endswith(".safetensors")) + assert st_file["metadata"]["format"] == "SafeTensor" + + # Images: the model-level thumbnail leads (the card preview derives from + # images[0]); paired entries display their site-hosted thumbnail, never + # the ephemeral imgdiff originals. + images = result["images"] + assert len(images) == 2 + assert images[0]["url"] == "https://openmodeldb.info/thumbs/ultrasharp-sr.jpg" + assert images[0]["nsfwLevel"] == 1 + assert images[1]["url"] == "https://openmodeldb.info/thumbs/small/t.jpg" + assert images[1]["meta"] == {"caption": "comparison"} + # The thumbnail IS the display URL here, so no separate thumbnailUrl. + assert "thumbnailUrl" not in images[1] + + # OpenModelDB-native provenance block + omdb = result["openmodeldb"] + assert omdb["id"] == "4x-UltraSharp" + assert omdb["url"] == "https://openmodeldb.info/models/4x-UltraSharp" + assert omdb["authors"] == ["kim"] + assert omdb["architecture"] == "esrgan" + assert omdb["architectureName"] == "ESRGAN" + assert omdb["scale"] == 4 + + +@pytest.mark.asyncio +async def test_lookup_supports_list_authors_and_standalone_images(make_client, dumps): + downloader = DummyDownloader(payloads=_payloads_by_url(dumps)) + client = make_client(downloader) + + result, error = await client.get_model_by_hash("c" * 64) + + assert error is None + assert result["creator"]["username"] == "Alice, Bob" + assert result["openmodeldb"]["authors"] == ["alice", "bob"] + assert result["images"][0]["url"] == "https://img.example.com/dejpg.png" + assert result["model"]["tags"] == ["Restoration"] + + +@pytest.mark.asyncio +async def test_lookup_never_emits_ephemeral_imgdiff_urls(make_client, dumps): + downloader = DummyDownloader(payloads=_payloads_by_url(dumps)) + client = make_client(downloader) + + result, error = await client.get_model_by_hash("d" * 64) + + assert error is None + images = result["images"] + # Model thumbnail first, then the one pair that has a thumbnail; the + # thumbnail-less imgdiff pair has no displayable asset and is dropped. + assert [img["url"] for img in images] == [ + "https://openmodeldb.info/thumbs/sonic-sr.jpg", + "https://openmodeldb.info/thumbs/small/sonic1.jpg", + ] + assert all("imgdiff.net" not in img["url"] for img in images) + # The original viewer URL survives as reference-only metadata. + assert images[1]["meta"]["comparisonUrl"] == ( + "https://imgdiff.net/api/image.php?id=abc&image=2172" + ) + assert images[1]["meta"]["caption"] == "Tails" + + +@pytest.mark.asyncio +async def test_lookup_derives_mediafire_mid_path_filename(make_client, dumps): + downloader = DummyDownloader(payloads=_payloads_by_url(dumps)) + client = make_client(downloader) + + result, error = await client.get_model_by_hash("d" * 64) + + assert error is None + # The basename of the mediafire URL is "file"; the real name sits mid-path. + assert result["files"][0]["name"] == "90s_Sonic_2x.pth" + # No direct-bytes mirror: the HTML-gateway URL stays as the reference. + assert result["files"][0]["downloadUrl"].startswith("https://www.mediafire.com/") + + +@pytest.mark.asyncio +async def test_lookup_synthesizes_filename_for_folder_links(make_client, dumps): + downloader = DummyDownloader(payloads=_payloads_by_url(dumps)) + client = make_client(downloader) + + result, error = await client.get_model_by_hash("e" * 64) + + assert error is None + # A mega.nz folder link has no filename; the catalogue type is authoritative. + assert result["files"][0]["name"] == "mega-only.pth" + + +@pytest.mark.asyncio +async def test_lookup_is_case_insensitive_and_reports_misses(make_client, dumps): + downloader = DummyDownloader(payloads=_payloads_by_url(dumps)) + client = make_client(downloader) + + result, error = await client.get_model_by_hash(MODEL_SHA.upper()) + assert error is None and result is not None + + result, error = await client.get_model_by_hash("f" * 64) + assert result is None + assert error == "Model not found" + + +@pytest.mark.asyncio +async def test_lookup_unavailable_without_cache_or_network(make_client): + downloader = DummyDownloader(fail=True) + client = make_client(downloader) + + result, error = await client.get_model_by_hash(MODEL_SHA) + + assert result is None + assert error == "OpenModelDB catalogue unavailable" + + +@pytest.mark.asyncio +async def test_lookup_survives_missing_optional_dumps(make_client, dumps): + # users/tags/architectures failing must degrade to raw ids, not break. + payloads = _payloads_by_url(dumps) + payloads[_dump_url("users")] = None + payloads[_dump_url("tags")] = None + payloads[_dump_url("architectures")] = None + downloader = DummyDownloader(payloads=payloads) + client = make_client(downloader) + + result, error = await client.get_model_by_hash(MODEL_SHA) + + assert error is None + assert result["creator"]["username"] == "kim" # raw id fallback + assert result["model"]["tags"] == ["general"] # raw id fallback + assert result["openmodeldb"]["architectureName"] == "esrgan" + + +# --------------------------------------------------------------------------- +# Disk cache, TTL, and conditional revalidation +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_disk_cache_written_and_reused_without_network(make_client, dumps): + payloads = _payloads_by_url(dumps) + downloader = DummyDownloader(payloads=payloads) + client = make_client(downloader) + + result, _ = await client.get_model_by_hash(MODEL_SHA) + assert result is not None + assert len(downloader.get_calls) == 4 # all four dumps fetched once + + # A second lookup hits the warm in-memory index: no additional HTTP. + result, _ = await client.get_model_by_hash(MODEL_SHA) + assert result is not None + assert len(downloader.get_calls) == 4 + assert len(downloader.head_calls) == 4 # HEAD probes from the initial fetch only + + # A fresh client instance (simulating a restart) within the TTL reads the + # disk cache without any network traffic. + offline = DummyDownloader(fail=True) + client2 = make_client(offline) + result, error = await client2.get_model_by_hash(MODEL_SHA) + assert error is None + assert result["name"] == "4x UltraSharp" + assert offline.get_calls == [] + assert offline.head_calls == [] + + +@pytest.mark.asyncio +async def test_expired_cache_revalidates_with_etag(make_client, dumps, tmp_path): + cache_dir = tmp_path / "omdb" + payloads = _payloads_by_url(dumps) + etags = {name: f'"etag-{name}"' for name in dumps} + headers = { + url: {"ETag": etags[name], "Last-Modified": "Wed, 01 Jan 2025 00:00:00 GMT"} + for name, url in ((n, _dump_url(n)) for n in dumps) + } + downloader = DummyDownloader(payloads=payloads, headers=headers) + client = make_client(downloader) + + await client.get_model_by_hash(MODEL_SHA) + assert len(downloader.get_calls) == 4 + + # Force the disk cache to look stale. + meta_path = cache_dir / "_meta.json" + meta = json.loads(meta_path.read_text()) + meta["fetched_at"] = time.time() - (CACHE_TTL_SECONDS + 60) + meta_path.write_text(json.dumps(meta)) + + # Unchanged etags -> cached bodies reused, no GET re-download. + downloader2 = DummyDownloader(payloads=None, headers=headers) + client2 = make_client(downloader2) + result, error = await client2.get_model_by_hash(MODEL_SHA) + assert error is None and result is not None + assert len(downloader2.head_calls) == 4 + assert downloader2.get_calls == [] + + # The TTL marker is refreshed even though bodies were reused. + meta = json.loads(meta_path.read_text()) + assert meta["fetched_at"] > time.time() - 60 + assert meta["etags"]["models"] == '"etag-models"' + + +@pytest.mark.asyncio +async def test_expired_cache_refetches_changed_dump(make_client, dumps, tmp_path): + cache_dir = tmp_path / "omdb" + payloads = _payloads_by_url(dumps) + old_etags = { + _dump_url(name): {"ETag": f'"old-{name}"'} for name in dumps + } + downloader = DummyDownloader(payloads=payloads, headers=old_etags) + client = make_client(downloader) + await client.get_model_by_hash(MODEL_SHA) + + meta_path = cache_dir / "_meta.json" + meta = json.loads(meta_path.read_text()) + meta["fetched_at"] = time.time() - (CACHE_TTL_SECONDS + 60) + meta_path.write_text(json.dumps(meta)) + + # Server reports a new etag for models.json only -> only it is re-downloaded. + new_headers = dict(old_etags) + new_headers[_dump_url("models")] = {"ETag": '"new-models"'} + downloader2 = DummyDownloader(payloads=payloads, headers=new_headers) + client2 = make_client(downloader2) + result, error = await client2.get_model_by_hash(MODEL_SHA) + assert error is None and result is not None + assert downloader2.get_calls == [_dump_url("models")] + + meta = json.loads(meta_path.read_text()) + assert meta["etags"]["models"] == '"new-models"' + assert meta["etags"]["users"] == '"old-users"' + + +@pytest.mark.asyncio +async def test_network_failure_falls_back_to_stale_cache(make_client, dumps, tmp_path): + cache_dir = tmp_path / "omdb" + payloads = _payloads_by_url(dumps) + downloader = DummyDownloader(payloads=payloads) + client = make_client(downloader) + await client.get_model_by_hash(MODEL_SHA) + + meta_path = cache_dir / "_meta.json" + meta = json.loads(meta_path.read_text()) + meta["fetched_at"] = time.time() - (CACHE_TTL_SECONDS + 60) + meta_path.write_text(json.dumps(meta)) + + offline = DummyDownloader(fail=True) + client2 = make_client(offline) + result, error = await client2.get_model_by_hash(MODEL_SHA) + assert error is None + assert result["name"] == "4x UltraSharp" + + +# --------------------------------------------------------------------------- +# Provider wrapper +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_provider_delegates_and_declares_unsupported_surface(make_client, dumps): + downloader = DummyDownloader(payloads=_payloads_by_url(dumps)) + client = make_client(downloader) + provider = OpenModelDBModelMetadataProvider(client) + + result, error = await provider.get_model_by_hash(MODEL_SHA) + assert error is None and result is not None + + assert await provider.get_model_versions("4x-UltraSharp") is None + assert await provider.get_model_version(model_id=1) is None + version, version_error = await provider.get_model_version_info("123") + assert version is None and version_error == "Model not found" + assert await provider.get_user_models("kim") is None + + +@pytest.mark.asyncio +async def test_fallback_chain_continues_past_openmodeldb_miss(make_client, dumps): + downloader = DummyDownloader(payloads=_payloads_by_url(dumps)) + client = make_client(downloader) + openmodeldb = OpenModelDBModelMetadataProvider(client) + + class _HitProvider(ModelMetadataProvider): + async def get_model_by_hash(self, model_hash): + return {"source": "elsewhere", "model": {"name": "Hit"}}, None + + async def get_model_versions(self, model_id): + return None + + async def get_model_version(self, model_id=None, version_id=None): + return None + + async def get_model_version_info(self, version_id): + return None, None + + async def get_user_models(self, username, cursor=None): + return None + + fallback = FallbackMetadataProvider( + [("openmodeldb_api", openmodeldb), ("other", _HitProvider())] + ) + + result, error = await fallback.get_model_by_hash("f" * 64) + assert result is not None + assert result["model"]["name"] == "Hit" + + +# --------------------------------------------------------------------------- +# FallbackMetadataProvider.excluding + sync-service sub_type gating +# --------------------------------------------------------------------------- + + +class _RecordingProvider(ModelMetadataProvider): + def __init__(self, label: str, result: Optional[Dict[str, Any]] = None) -> None: + self.label = label + self.result = result + self.calls: list = [] + + async def get_model_by_hash(self, model_hash: str): + self.calls.append(model_hash) + return self.result, None if self.result else "Model not found" + + async def get_model_versions(self, model_id): + return None + + async def get_model_version(self, model_id=None, version_id=None): + return None + + async def get_model_version_info(self, version_id): + return None, None + + async def get_user_models(self, username, cursor=None): + return None + + +def test_fallback_excluding_drops_named_providers(): + civitai = _RecordingProvider("civitai_api") + openmodeldb = _RecordingProvider("openmodeldb_api") + fallback = FallbackMetadataProvider( + [("civitai_api", civitai), ("openmodeldb_api", openmodeldb)] + ) + + filtered = fallback.excluding({"openmodeldb_api"}) + assert filtered is not fallback + assert filtered._provider_labels == ["civitai_api"] + # Original chain is untouched. + assert fallback._provider_labels == ["civitai_api", "openmodeldb_api"] + # Empty/no-op exclusions return the same instance. + assert fallback.excluding(set()) is fallback + + +def _build_sync_service(default_provider, provider_selector): + metadata_manager = SimpleNamespace(save_metadata=AsyncMock()) + preview_service = SimpleNamespace(ensure_preview_for_metadata=AsyncMock()) + settings = SimpleNamespace(get=lambda key, default=None: default) + return MetadataSyncService( + metadata_manager=metadata_manager, + preview_service=preview_service, + settings=settings, + default_metadata_provider_factory=AsyncMock(return_value=default_provider), + metadata_provider_selector=provider_selector, + ) + + +@pytest.mark.asyncio +async def test_sync_gates_openmodeldb_out_for_non_upscalers(tmp_path): + civitai = _RecordingProvider("civitai_api") + openmodeldb = _RecordingProvider( + "openmodeldb_api", result={"source": "openmodeldb", "model": {"name": "OMDB"}} + ) + fallback = FallbackMetadataProvider( + [("civitai_api", civitai), ("openmodeldb_api", openmodeldb)] + ) + service = _build_sync_service(fallback, AsyncMock()) + + model_data: Dict[str, Any] = { + "model_name": "lora", + "file_path": str(tmp_path / "lora.safetensors"), + # No sub_type (a LoRA): openmodeldb must not be consulted. + } + ok, _ = await service.fetch_and_update_model( + sha256="abc", + file_path=model_data["file_path"], + model_data=model_data, + update_cache_func=AsyncMock(return_value=True), + ) + assert ok is False # nothing found anywhere + assert civitai.calls == ["abc"] + assert openmodeldb.calls == [] + + +@pytest.mark.asyncio +async def test_sync_consults_openmodeldb_for_upscalers(tmp_path): + civitai = _RecordingProvider("civitai_api") + openmodeldb = _RecordingProvider( + "openmodeldb_api", + result={ + "source": "openmodeldb", + "name": "4x UltraSharp", + "model": {"name": "4x UltraSharp", "type": "Upscaler", "tags": []}, + "files": [{"name": "x.pth", "hashes": {"SHA256": "ABC"}}], + "images": [], + }, + ) + fallback = FallbackMetadataProvider( + [("civitai_api", civitai), ("openmodeldb_api", openmodeldb)] + ) + service = _build_sync_service(fallback, AsyncMock()) + + model_data: Dict[str, Any] = { + "model_name": "upscaler", + "file_path": str(tmp_path / "upscaler.pth"), + "sub_type": "upscaler", + } + ok, error = await service.fetch_and_update_model( + sha256="abc", + file_path=model_data["file_path"], + model_data=model_data, + update_cache_func=AsyncMock(return_value=True), + ) + assert ok is True, error + assert civitai.calls == ["abc"] + assert openmodeldb.calls == ["abc"] + assert model_data["metadata_source"] == "openmodeldb" + + +@pytest.mark.asyncio +async def test_sync_deleted_upscaler_still_reaches_openmodeldb(tmp_path): + openmodeldb = _RecordingProvider( + "openmodeldb_api", + result={ + "source": "openmodeldb", + "name": "Recovered", + "model": {"name": "Recovered", "type": "Upscaler", "tags": []}, + "files": [], + "images": [], + }, + ) + + async def selector(name): + if name == "openmodeldb_api": + return openmodeldb + raise ValueError(f"Provider '{name}' is not registered") + + service = _build_sync_service( + SimpleNamespace(get_model_by_hash=AsyncMock(return_value=(None, "Model not found"))), + AsyncMock(side_effect=selector), + ) + + model_data: Dict[str, Any] = { + "model_name": "deleted upscaler", + "file_path": str(tmp_path / "deleted.pth"), + "sub_type": "upscaler", + "civitai_deleted": True, + "metadata_source": "openmodeldb", + } + ok, error = await service.fetch_and_update_model( + sha256="abc", + file_path=model_data["file_path"], + model_data=model_data, + update_cache_func=AsyncMock(return_value=True), + ) + assert ok is True, error + assert openmodeldb.calls == ["abc"] + assert model_data["metadata_source"] == "openmodeldb" + + +@pytest.mark.asyncio +async def test_sync_openmodeldb_sourced_model_prefers_openmodeldb_provider(tmp_path): + """A model downloaded from OpenModelDB refreshes against its catalogue.""" + civitai = _RecordingProvider("civitai_api") + openmodeldb = _RecordingProvider( + "openmodeldb_api", + result={ + "source": "openmodeldb", + "name": "4x UltraSharp", + "model": {"name": "4x UltraSharp", "type": "Upscaler", "tags": []}, + "files": [{"name": "x.pth", "hashes": {"SHA256": "ABC"}}], + "images": [], + }, + ) + + async def selector(name): + return {"civitai_api": civitai, "openmodeldb_api": openmodeldb}[name] + + service = _build_sync_service( + SimpleNamespace(get_model_by_hash=AsyncMock(return_value=(None, "Model not found"))), + AsyncMock(side_effect=selector), + ) + + model_data: Dict[str, Any] = { + "model_name": "upscaler", + "file_path": str(tmp_path / "upscaler.pth"), + "sub_type": "upscaler", + "source_platform": "openmodeldb", + "source_url": "https://openmodeldb.info/models/4x-UltraSharp", + } + ok, error = await service.fetch_and_update_model( + sha256="abc", + file_path=model_data["file_path"], + model_data=model_data, + update_cache_func=AsyncMock(return_value=True), + ) + assert ok is True, error + # The source's own catalogue answers first; CivitAI is not needed. + assert openmodeldb.calls == ["abc"] + assert civitai.calls == [] + assert model_data["metadata_source"] == "openmodeldb" + # A "not found" from the source provider must not read as civitai-deleted. + assert model_data.get("civitai_deleted") is not True + + +@pytest.mark.asyncio +async def test_sync_openmodeldb_sourced_model_falls_back_to_civitai(tmp_path): + """When the OpenModelDB catalogue has no record, CivitAI is still tried.""" + civitai = _RecordingProvider( + "civitai_api", + result={ + "source": "civitai_api", + "name": "Also on CivitAI", + "model": {"name": "Also on CivitAI", "type": "Upscaler", "tags": []}, + "files": [{"name": "x.pth", "hashes": {"SHA256": "ABC"}}], + "images": [], + }, + ) + openmodeldb = _RecordingProvider("openmodeldb_api") + + async def selector(name): + return {"civitai_api": civitai, "openmodeldb_api": openmodeldb}[name] + + service = _build_sync_service( + SimpleNamespace(get_model_by_hash=AsyncMock(return_value=(None, "Model not found"))), + AsyncMock(side_effect=selector), + ) + + model_data: Dict[str, Any] = { + "model_name": "upscaler", + "file_path": str(tmp_path / "upscaler.pth"), + "sub_type": "upscaler", + "source_platform": "openmodeldb", + "source_url": "https://openmodeldb.info/models/4x-UltraSharp", + } + ok, error = await service.fetch_and_update_model( + sha256="abc", + file_path=model_data["file_path"], + model_data=model_data, + update_cache_func=AsyncMock(return_value=True), + ) + assert ok is True, error + assert openmodeldb.calls == ["abc"] + assert civitai.calls == ["abc"] + assert model_data["metadata_source"] == "civitai_api" + + +@pytest.mark.asyncio +async def test_sync_huggingface_sourced_model_stays_civitai_only(tmp_path): + """Sources without their own provider keep the CivitAI-only behaviour.""" + civitai = _RecordingProvider("civitai_api") + + async def selector(name): + return {"civitai_api": civitai}[name] + + service = _build_sync_service( + SimpleNamespace(get_model_by_hash=AsyncMock(return_value=(None, "Model not found"))), + AsyncMock(side_effect=selector), + ) + + model_data: Dict[str, Any] = { + "model_name": "lora", + "file_path": str(tmp_path / "lora.safetensors"), + "source_platform": "huggingface", + "source_url": "https://huggingface.co/user/repo", + } + ok, _ = await service.fetch_and_update_model( + sha256="abc", + file_path=model_data["file_path"], + model_data=model_data, + update_cache_func=AsyncMock(return_value=True), + ) + assert ok is False + assert civitai.calls == ["abc"] + # A "Model not found" from the named provider does not mark deletion. + assert model_data.get("civitai_deleted") is not True diff --git a/tests/services/test_openmodeldb_source.py b/tests/services/test_openmodeldb_source.py new file mode 100644 index 00000000..8f294d95 --- /dev/null +++ b/tests/services/test_openmodeldb_source.py @@ -0,0 +1,332 @@ +"""Tests for the OpenModelDB model source (URL parsing, downloads, card context).""" + +from __future__ import annotations + +from typing import Any, Dict +from unittest.mock import AsyncMock + +import pytest + +from py.services.model_sources import ( + ModelSourceError, + detect_source, + get_download_source, + source_group_key, +) +from py.services.model_sources.openmodeldb import OpenModelDBSource +from py.services.openmodeldb_client import OpenModelDBClient + +MODEL_SHA = "a" * 64 +ALT_SHA = "b" * 64 + + +def _dumps() -> Dict[str, Dict[str, Any]]: + return { + "models": { + "4x-UltraSharp": { + "name": "4x UltraSharp", + "author": "kim", + "license": "MIT", + "tags": ["general"], + "description": "Sharp general-purpose upscaler", + "date": "2023-01-02", + "architecture": "esrgan", + "scale": 4, + "resources": [ + { + "platform": "pytorch", + "type": "pth", + "size": 67_000_000, + "sha256": MODEL_SHA, + "urls": [ + "https://files.example.com/4x-UltraSharp.pth", + "https://mega.nz/file/mirror", + ], + }, + { + "platform": "pytorch", + "type": "safetensors", + "size": 67_100_000, + "sha256": ALT_SHA, + "urls": ["https://files.example.com/4x-UltraSharp.safetensors"], + }, + ], + "images": [ + { + "type": "paired", + "LR": "https://img.example.com/lr.jpg", + "SR": "https://img.example.com/sr.jpg", + } + ], + }, + "onnx-only": { + "name": "ONNX Only", + "author": "kim", + "architecture": "esrgan", + "scale": 2, + "resources": [ + { + "platform": "onnx", + "type": "onnx", + "size": 10_000, + "sha256": "c" * 64, + "urls": ["https://files.example.com/onnx-only.onnx"], + } + ], + }, + "2x-90s-Sonic-LG": { + "name": "90s Sonic 2x (Large)", + "author": "kim", + "architecture": "esrgan", + "scale": 2, + "resources": [ + { + "platform": "pytorch", + "type": "pth", + "size": 9_443_322, + "sha256": "d" * 64, + "urls": [ + "https://www.mediafire.com/file/e57pd40qph8nak2/90s_Sonic_2x.pth/file" + ], + } + ], + }, + "mixed-mirror": { + "name": "Mixed Mirror", + "author": "kim", + "architecture": "esrgan", + "scale": 4, + "resources": [ + { + "platform": "pytorch", + "type": "pth", + "size": 60_000_000, + "sha256": "e" * 64, + # Gateway host first, direct mirror second. + "urls": [ + "https://mega.nz/folder/qZRBmaIY#nIG8KyWFcGNTuMX_XNbJ_g", + "https://files.example.com/Mixed-Mirror.pth", + ], + } + ], + }, + }, + "users": {"kim": {"name": "Kim"}}, + "tags": {"general": {"name": "General Purpose"}}, + "architectures": {"esrgan": {"name": "ESRGAN"}}, + } + + +@pytest.fixture +def client(tmp_path, monkeypatch) -> OpenModelDBClient: + """A catalogue-loaded client, installed as the singleton for the source.""" + seeded = OpenModelDBClient(cache_dir=str(tmp_path / "omdb")) + seeded._install_payloads(_dumps()) + monkeypatch.setattr( + OpenModelDBClient, "get_instance", AsyncMock(return_value=seeded) + ) + return seeded + + +# --------------------------------------------------------------------------- +# Identity and URL parsing +# --------------------------------------------------------------------------- + + +def test_url_parsing_lenient_and_strict(): + source = OpenModelDBSource() + assert source.parse("https://openmodeldb.info/models/4x-UltraSharp") == "4x-UltraSharp" + assert ( + source.parse("https://openmodeldb.info/models/4x-UltraSharp", strict=True) + == "4x-UltraSharp" + ) + assert ( + source.parse("https://openmodeldb.info/models/4x-UltraSharp/", strict=True) + == "4x-UltraSharp" + ) + # Query strings are tolerated only in lenient mode. + assert source.parse("https://openmodeldb.info/models/4x-UltraSharp?x=1") == "4x-UltraSharp" + assert source.parse("https://openmodeldb.info/models/4x-UltraSharp?x=1", strict=True) is None + assert source.parse("https://openmodeldb.info/") is None + + +def test_detect_source_and_group_key(): + ref = detect_source("https://openmodeldb.info/models/4x-UltraSharp") + assert ref is not None + assert ref.platform == "openmodeldb" + assert ref.source_id == "4x-UltraSharp" + assert ref.url == "https://openmodeldb.info/models/4x-UltraSharp" + + # The model id IS the published-model identity, so downloads group by it. + assert ( + source_group_key( + {"source_url": ref.url, "source_platform": "openmodeldb"} + ) + == "omdb:4x-UltraSharp" + ) + + +def test_flat_source_id_validation_and_default_paths(): + source = get_download_source("openmodeldb") + assert source is not None + assert source.is_valid_source_id("4x-UltraSharp") is True + assert source.is_valid_source_id("owner/name") is False + assert source.is_valid_source_id("../escape") is False + assert source.is_valid_source_id("") is False + # Flat catalogue: no owner/repo namespace under the default directory. + assert source.default_subdir_parts("4x-UltraSharp") == ("openmodeldb",) + + +def test_capabilities(): + source = OpenModelDBSource() + assert source.supports_download is True + assert source.supports_enrichment is True + assert source.example_source_id + assert source.canonical_url(source.example_source_id).startswith( + "https://openmodeldb.info/models/" + ) + + +# --------------------------------------------------------------------------- +# File listing and download URL resolution +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_list_files_returns_pytorch_resources_largest_first(client): + files = await OpenModelDBSource().list_files("4x-UltraSharp") + assert [f["filename"] for f in files] == [ + "4x-UltraSharp.safetensors", + "4x-UltraSharp.pth", + ] + assert files[0]["size"] == 67_100_000 + + +@pytest.mark.asyncio +async def test_list_files_skips_non_pytorch_resources(client): + # `.onnx` is not a loadable weight format, so an onnx-only model has an + # empty download list rather than offering an unusable file. + assert await OpenModelDBSource().list_files("onnx-only") == [] + + +@pytest.mark.asyncio +async def test_list_files_raises_manual_download_hint_when_mirror_only(client): + # Every mirror is an HTML gateway: an empty list would surface in the UI + # as "no model files found", so the source reports the actionable cause. + with pytest.raises(ModelSourceError) as excinfo: + await OpenModelDBSource().list_files("2x-90s-Sonic-LG") + assert excinfo.value.status == 400 + assert "manual" in str(excinfo.value).lower() + assert "openmodeldb.info/models/2x-90s-Sonic-LG" in str(excinfo.value) + + +@pytest.mark.asyncio +async def test_list_files_prefers_direct_mirror_over_gateway(client): + files = await OpenModelDBSource().list_files("mixed-mirror") + assert files == [{"filename": "Mixed-Mirror.pth", "size": 60_000_000}] + + +@pytest.mark.asyncio +async def test_resolve_download_url_rejects_html_gateway_mirror(client): + with pytest.raises(ModelSourceError) as excinfo: + await OpenModelDBSource().resolve_download_url( + "2x-90s-Sonic-LG", "90s_Sonic_2x.pth" + ) + assert excinfo.value.status == 400 + message = str(excinfo.value) + assert "mediafire.com" in message + assert "openmodeldb.info/models/2x-90s-Sonic-LG" in message + + +@pytest.mark.asyncio +async def test_resolve_download_url_uses_direct_mirror(client): + url = await OpenModelDBSource().resolve_download_url( + "mixed-mirror", "Mixed-Mirror.pth" + ) + assert url == "https://files.example.com/Mixed-Mirror.pth" + + +@pytest.mark.asyncio +async def test_list_files_unknown_model_raises_404(client): + with pytest.raises(ModelSourceError) as excinfo: + await OpenModelDBSource().list_files("nope") + assert excinfo.value.status == 404 + + +@pytest.mark.asyncio +async def test_resolve_download_url_uses_primary_url_only(client): + source = OpenModelDBSource() + url = await source.resolve_download_url("4x-UltraSharp", "4x-UltraSharp.pth") + # Mirrors (mega.nz & co.) need site-specific handling and are never used. + assert url == "https://files.example.com/4x-UltraSharp.pth" + assert ( + await source.resolve_download_url("4x-UltraSharp", "4x-UltraSharp.safetensors") + == "https://files.example.com/4x-UltraSharp.safetensors" + ) + + +@pytest.mark.asyncio +async def test_resolve_download_url_unknown_file_raises_404(client): + with pytest.raises(ModelSourceError) as excinfo: + await OpenModelDBSource().resolve_download_url("4x-UltraSharp", "nope.pth") + assert excinfo.value.status == 404 + + +@pytest.mark.asyncio +async def test_unavailable_catalogue_raises_502(tmp_path, monkeypatch): + broken = OpenModelDBClient(cache_dir=str(tmp_path / "omdb")) + monkeypatch.setattr(broken, "_ensure_loaded", AsyncMock(return_value=False)) + monkeypatch.setattr( + OpenModelDBClient, "get_instance", AsyncMock(return_value=broken) + ) + + with pytest.raises(ModelSourceError) as excinfo: + await OpenModelDBSource().list_files("4x-UltraSharp") + assert excinfo.value.status == 502 + + with pytest.raises(ModelSourceError) as excinfo: + await OpenModelDBSource().resolve_download_url("4x-UltraSharp", "x.pth") + assert excinfo.value.status == 502 + + +# --------------------------------------------------------------------------- +# Card context (download-time hydration) +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_fetch_model_card_context(client): + context = await OpenModelDBSource().fetch_model_card_context( + "4x-UltraSharp", "4x-UltraSharp.pth", sha256=MODEL_SHA + ) + assert context.description == "Sharp general-purpose upscaler" + assert context.model_name == "4x UltraSharp" + assert context.license == "MIT" + assert context.model_type == "Upscaler" + assert context.official_tags == ["General Purpose"] + # Paired previews show the upscaled (SR) result. + assert context.example_images == ["https://img.example.com/sr.jpg"] + assert context.source_model_id == "4x-UltraSharp" + assert context.is_empty() is False + + +@pytest.mark.asyncio +async def test_fetch_model_card_context_unknown_model_is_empty(client): + context = await OpenModelDBSource().fetch_model_card_context("nope") + assert context.is_empty() is True + + +@pytest.mark.asyncio +async def test_fetch_model_card_context_survives_catalogue_failure( + tmp_path, monkeypatch +): + broken = OpenModelDBClient(cache_dir=str(tmp_path / "omdb")) + monkeypatch.setattr( + broken, "_ensure_loaded", AsyncMock(side_effect=RuntimeError("boom")) + ) + monkeypatch.setattr( + OpenModelDBClient, "get_instance", AsyncMock(return_value=broken) + ) + # Never raises: a lookup fault reads as "the site had nothing extra". + context = await OpenModelDBSource().fetch_model_card_context("4x-UltraSharp") + assert context.is_empty() is True