feat(metadata): add OpenModelDB metadata provider and model source for upscalers

Add OpenModelDB (openmodeldb.info) as a metadata and download source for
the existing upscaler model type.

Metadata:
- New OpenModelDBClient: fetches the site's bulk JSON dumps, caches them
  on disk (24h TTL + ETag revalidation), and builds a local sha256 index
- New OpenModelDBModelMetadataProvider adapts catalogue entries to the
  CivitAI-shaped version dict contract; registered in the fallback chain
  behind the enable_openmodeldb_api setting (default on), gated to the
  upscaler sub-type so other model types never trigger the dump download
- Persisted provenance uses metadata_source "openmodeldb" plus a nested
  openmodeldb block (page URL, architecture, scale, license)

Images: paired-image LR/SR URLs are ephemeral imgdiff.net sessions, so
displayable images come from the site-hosted auto-generated thumbnails
(model-level cover leads images[], per-image thumbs for the rest); the
original comparison URL is kept in meta.comparisonUrl.

Downloads:
- New OpenModelDBSource (flat model ids, omdb: group prefix) with
  resource filename derivation that recovers names hidden mid-path
  (mediafire) or synthesizes {id}.{type} for folder links
- HTML-gateway mirrors (mediafire/mega/drive) are rejected with a clear
  manual-download hint instead of silently saving an HTML page as .pth
- ModelSource base gains is_valid_source_id / default_subdir_parts /
  resolve_download_url hooks so flat-id sources need no platform branches

UI: "View on OpenModelDB" link in the model modal (downloaded and
hash-enriched models), settings toggle next to the CivArchive one.
This commit is contained in:
Will Miao
2026-10-03 21:13:07 +08:00
parent 515469054c
commit 35f1ced41a
39 changed files with 2787 additions and 37 deletions
+51
View File
@@ -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"]
+3 -1
View File
@@ -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
+862
View File
@@ -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
+332
View File
@@ -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