mirror of
https://github.com/willmiao/ComfyUI-Lora-Manager.git
synced 2026-10-04 00:55:32 -03:00
feat(metadata): add OpenModelDB metadata provider and model source for upscalers
Add OpenModelDB (openmodeldb.info) as a metadata and download source for
the existing upscaler model type.
Metadata:
- New OpenModelDBClient: fetches the site's bulk JSON dumps, caches them
on disk (24h TTL + ETag revalidation), and builds a local sha256 index
- New OpenModelDBModelMetadataProvider adapts catalogue entries to the
CivitAI-shaped version dict contract; registered in the fallback chain
behind the enable_openmodeldb_api setting (default on), gated to the
upscaler sub-type so other model types never trigger the dump download
- Persisted provenance uses metadata_source "openmodeldb" plus a nested
openmodeldb block (page URL, architecture, scale, license)
Images: paired-image LR/SR URLs are ephemeral imgdiff.net sessions, so
displayable images come from the site-hosted auto-generated thumbnails
(model-level cover leads images[], per-image thumbs for the rest); the
original comparison URL is kept in meta.comparisonUrl.
Downloads:
- New OpenModelDBSource (flat model ids, omdb: group prefix) with
resource filename derivation that recovers names hidden mid-path
(mediafire) or synthesizes {id}.{type} for folder links
- HTML-gateway mirrors (mediafire/mega/drive) are rejected with a clear
manual-download hint instead of silently saving an HTML page as .pth
- ModelSource base gains is_valid_source_id / default_subdir_parts /
resolve_download_url hooks so flat-id sources need no platform branches
UI: "View on OpenModelDB" link in the model modal (downloaded and
hash-enriched models), settings toggle next to the CivArchive one.
This commit is contained in:
@@ -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');
|
||||
});
|
||||
});
|
||||
|
||||
@@ -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(() => {});
|
||||
|
||||
@@ -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')
|
||||
|
||||
@@ -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}
|
||||
]
|
||||
|
||||
@@ -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"]
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
@@ -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
|
||||
Reference in New Issue
Block a user