mirror of
https://github.com/willmiao/ComfyUI-Lora-Manager.git
synced 2026-10-03 16:45:33 -03:00
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.
863 lines
31 KiB
Python
863 lines
31 KiB
Python
"""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
|