From 196c8ffc3e06df65ffd8255941d4e803e0454d1a Mon Sep 17 00:00:00 2001 From: Will Miao Date: Sat, 8 Aug 2026 22:12:47 +0800 Subject: [PATCH] feat(recipes): match civitai image hash sections against local hash cache --- py/recipes/parsers/civitai_image.py | 153 ++++++--- tests/services/test_civitai_image_parser.py | 328 ++++++++++++++++++++ 2 files changed, 435 insertions(+), 46 deletions(-) diff --git a/py/recipes/parsers/civitai_image.py b/py/recipes/parsers/civitai_image.py index ad9167c1..dab68ef9 100644 --- a/py/recipes/parsers/civitai_image.py +++ b/py/recipes/parsers/civitai_image.py @@ -4,7 +4,7 @@ import json import logging from typing import Dict, Any, Union from ..base import RecipeMetadataParser -from ..constants import GEN_PARAM_KEYS +from ..constants import GEN_PARAM_KEYS, VALID_LORA_TYPES from ...services.metadata_service import get_default_metadata_provider from ...config import config @@ -216,7 +216,8 @@ class CivitaiApiMetadataParser(RecipeMetadataParser): # Try to look up base model from the checkpoint hash cp_hash = checkpoint_entry.get("hash") if cp_hash and metadata_provider: - local_cached = local_cache.get(cp_hash) if local_cache else None + # local_cache keys are stored lowercase + local_cached = local_cache.get(cp_hash.lower()) if local_cache else None if local_cached: self._populate_entry_from_cache( checkpoint_entry, local_cached @@ -294,8 +295,15 @@ class CivitaiApiMetadataParser(RecipeMetadataParser): # Try to get info from Civitai if hash is available if lora_hash and metadata_provider: - local_cached = local_cache.get(lora_hash) if local_cache else None + # local_cache keys are stored lowercase + local_cached = local_cache.get(lora_hash.lower()) if local_cache else None if local_cached: + cached_type = self._cache_item_model_type(local_cached) + if cached_type and cached_type not in VALID_LORA_TYPES: + logger.debug( + f"Skipping non-LoRA cache item for hash {lora_hash}" + ) + continue self._populate_entry_from_cache( lora_entry, local_cached ) @@ -304,6 +312,12 @@ class CivitaiApiMetadataParser(RecipeMetadataParser): added_loras[str(lora_entry["id"])] = len( result["loras"] ) + # Mirror base.py:150-151 counts for API-path loras + bm = local_cached.get("base_model") or "" + if bm: + base_model_counts[bm] = base_model_counts.get( + bm, 0 + ) + 1 else: try: civitai_info = ( @@ -649,30 +663,47 @@ class CivitaiApiMetadataParser(RecipeMetadataParser): } if metadata_provider: - try: - civitai_info = await metadata_provider.get_model_by_hash( - lora_hash - ) - - populated_entry = await self.populate_lora_from_civitai( - lora_entry, - civitai_info, - recipe_scanner, - base_model_counts, - lora_hash, - ) - - if populated_entry is None: + # local_cache keys are stored lowercase + local_cached = local_cache.get(lora_hash.lower()) if local_cache else None + if local_cached: + cached_type = self._cache_item_model_type(local_cached) + if cached_type and cached_type not in VALID_LORA_TYPES: + logger.debug( + f"Skipping non-LoRA cache item for hash {lora_hash}" + ) continue - - lora_entry = populated_entry - + self._populate_entry_from_cache(lora_entry, local_cached) + # Mirror base.py:150-151 counts for API-path loras + bm = local_cached.get("base_model") or "" + if bm: + base_model_counts[bm] = base_model_counts.get(bm, 0) + 1 if "id" in lora_entry and lora_entry["id"]: added_loras[str(lora_entry["id"])] = len(result["loras"]) - except Exception as e: - logger.error( - f"Error fetching Civitai info for LoRA hash {lora_hash}: {e}" - ) + else: + try: + civitai_info = await metadata_provider.get_model_by_hash( + lora_hash + ) + + populated_entry = await self.populate_lora_from_civitai( + lora_entry, + civitai_info, + recipe_scanner, + base_model_counts, + lora_hash, + ) + + if populated_entry is None: + continue + + lora_entry = populated_entry + + if "id" in lora_entry and lora_entry["id"]: + added_loras[str(lora_entry["id"])] = len(result["loras"]) + except Exception as e: + logger.error( + f"Error fetching Civitai info for LoRA hash {lora_hash}: {e}" + ) added_loras[lora_hash] = len(result["loras"]) result["loras"].append(lora_entry) @@ -711,32 +742,51 @@ class CivitaiApiMetadataParser(RecipeMetadataParser): # Try to get info from Civitai if hash is available if lora_entry["hash"] and metadata_provider: - try: - civitai_info = await metadata_provider.get_model_by_hash( - lora_hash - ) - - populated_entry = await self.populate_lora_from_civitai( - lora_entry, - civitai_info, - recipe_scanner, - base_model_counts, - lora_hash, - ) - - if populated_entry is None: + # local_cache keys are stored lowercase + local_cached = local_cache.get(lora_hash.lower()) if local_cache else None + if local_cached: + cached_type = self._cache_item_model_type(local_cached) + if cached_type and cached_type not in VALID_LORA_TYPES: + logger.debug( + f"Skipping non-LoRA cache item for hash {lora_hash}" + ) lora_index += 1 - continue # Skip invalid LoRA types - - lora_entry = populated_entry - + continue # Skip non-LoRA cache items + self._populate_entry_from_cache(lora_entry, local_cached) + # Mirror base.py:150-151 counts for API-path loras + bm = local_cached.get("base_model") or "" + if bm: + base_model_counts[bm] = base_model_counts.get(bm, 0) + 1 # If we have a version ID from Civitai, track it for deduplication if "id" in lora_entry and lora_entry["id"]: added_loras[str(lora_entry["id"])] = len(result["loras"]) - except Exception as e: - logger.error( - f"Error fetching Civitai info for LoRA hash {lora_entry['hash']}: {e}" - ) + else: + try: + civitai_info = await metadata_provider.get_model_by_hash( + lora_hash + ) + + populated_entry = await self.populate_lora_from_civitai( + lora_entry, + civitai_info, + recipe_scanner, + base_model_counts, + lora_hash, + ) + + if populated_entry is None: + lora_index += 1 + continue # Skip invalid LoRA types + + lora_entry = populated_entry + + # If we have a version ID from Civitai, track it for deduplication + if "id" in lora_entry and lora_entry["id"]: + added_loras[str(lora_entry["id"])] = len(result["loras"]) + except Exception as e: + logger.error( + f"Error fetching Civitai info for LoRA hash {lora_entry['hash']}: {e}" + ) # Track by hash if we have it if lora_hash: @@ -795,3 +845,14 @@ class CivitaiApiMetadataParser(RecipeMetadataParser): base_model = cache_item.get("base_model", "") if base_model: entry["baseModel"] = base_model + + @staticmethod + def _cache_item_model_type(cache_item: dict[str, Any]) -> str: + """Lowercased civitai.model.type of a cache item, or '' when unknown.""" + civ = cache_item.get("civitai") + if not isinstance(civ, dict): + return "" + model_info = civ.get("model") + if not isinstance(model_info, dict): + return "" + return (model_info.get("type") or "").lower() diff --git a/tests/services/test_civitai_image_parser.py b/tests/services/test_civitai_image_parser.py index 629c10cb..7073c1dd 100644 --- a/tests/services/test_civitai_image_parser.py +++ b/tests/services/test_civitai_image_parser.py @@ -570,3 +570,331 @@ async def test_backfill_lora_cache_item_without_sha256_does_not_crash(): ) assert result["thumbnailUrl"].startswith("https://image.civitai.com/") + +class _RecordingProvider: + """Metadata provider stub that records get_model_by_hash calls. + + With raise_on_call=True any hash lookup fails the test loudly — used to + prove that local_cache hits skip the CivitAI API. Otherwise the result + is returned for the cache-miss path. + """ + + def __init__(self, result=None, raise_on_call=False): + self.hash_calls = [] + self._result = result + self._raise_on_call = raise_on_call + + async def get_model_by_hash(self, model_hash): + self.hash_calls.append(model_hash) + if self._raise_on_call: + raise AssertionError( + f"get_model_by_hash should not be called on local_cache hit, got {model_hash}" + ) + return self._result + + async def get_model_version_info(self, version_id): + return None, "Model not found" + + +def _cache_item(name="Local Style", base_model="SDXL 1.0", model_type="LORA"): + """Build a scanner cache item shaped like the local hash cache values.""" + return { + "file_path": f"/loras/{name.lower().replace(' ', '_')}.safetensors", + "file_name": name, + "sha256": "aabbccddeeff00112233445566778899aabbccddeeff00112233445566778899", + "autov3": "", + "preview_url": "/previews/style.png", + "base_model": base_model, + "civitai": { + "id": 300, + "modelId": 400, + "name": "v1", + "model": {"name": name, "type": model_type}, + }, + } + + +def _make_lora_info(base_model="SDXL 1.0"): + """Civitai response for a lora hash lookup on the cache-miss path.""" + return { + "id": 300, + "modelId": 400, + "model": {"name": "Style LoRA", "type": "lora"}, + "name": "v1", + "images": [{"url": "https://image.civitai.com/lora/original=true"}], + "baseModel": base_model, + "downloadUrl": "https://civitai.com/api/download/300", + "files": [ + { + "type": "Model", + "primary": True, + "sizeKB": 512, + "name": "style.safetensors", + "hashes": {"SHA256": "ff00112233445566778899aabbccddeeff00112233445566778899aabbccddee"}, + } + ], + } + + +async def _parse_with_cache(monkeypatch, provider, metadata, local_cache=None): + """Run parse_metadata with a fixed metadata provider and optional local_cache.""" + async def fake_metadata_provider(): + return provider + + monkeypatch.setattr( + "py.recipes.parsers.civitai_image.get_default_metadata_provider", + fake_metadata_provider, + ) + parser = CivitaiApiMetadataParser() + return await parser.parse_metadata(metadata, local_cache=local_cache) + + +@pytest.mark.asyncio +async def test_local_cache_hashes_section_populates_from_cache_and_skips_api(monkeypatch): + """Hashes-section lora whose hash is a local_cache key is populated from + the cache and the metadata provider is never consulted.""" + provider = _RecordingProvider(raise_on_call=True) + item = _cache_item(name="Local Style", base_model="SDXL 1.0") + local_cache = {"a1b2c3d4e5f6": item} + metadata = {"hashes": {"LORA:Local Style": "A1B2C3D4E5F6"}} + + result = await _parse_with_cache(monkeypatch, provider, metadata, local_cache=local_cache) + + assert provider.hash_calls == [] + assert len(result["loras"]) == 1 + lora = result["loras"][0] + assert lora["existsLocally"] is True + assert lora["localPath"] == item["file_path"] + assert lora["hash"] == item["sha256"] + assert lora["name"] == "Local Style" + + +@pytest.mark.asyncio +async def test_local_cache_lora_n_section_populates_from_cache_and_skips_api(monkeypatch): + """Lora_N section lora whose hash is a local_cache key is populated from + the cache and the metadata provider is never consulted.""" + provider = _RecordingProvider(raise_on_call=True) + item = _cache_item(name="Lora N Style", base_model="SDXL 1.0") + local_cache = {"abc123def456": item} + metadata = { + "Lora_0 Model hash": "ABC123DEF456", + "Lora_0 Model name": "Lora N Style", + "Lora_0 Strength model": 0.7, + } + + result = await _parse_with_cache(monkeypatch, provider, metadata, local_cache=local_cache) + + assert provider.hash_calls == [] + assert len(result["loras"]) == 1 + lora = result["loras"][0] + assert lora["existsLocally"] is True + assert lora["localPath"] == item["file_path"] + assert lora["weight"] == 0.7 + assert lora["hash"] == item["sha256"] + + +@pytest.mark.asyncio +async def test_local_cache_uppercase_hash_matches_lowercase_key(monkeypatch): + """Resources lora with an UPPERCASE hash still matches the lowercase key.""" + provider = _RecordingProvider(raise_on_call=True) + sha256 = "aabbccddeeff00112233445566778899aabbccddeeff00112233445566778899" + item = _cache_item(name="Upper Case", base_model="SDXL 1.0") + local_cache = {sha256: item} + metadata = { + "resources": [ + {"hash": sha256.upper(), "name": "Upper Case", "type": "lora", "weight": 0.5}, + ], + } + + result = await _parse_with_cache(monkeypatch, provider, metadata, local_cache=local_cache) + + assert provider.hash_calls == [] + assert len(result["loras"]) == 1 + assert result["loras"][0]["existsLocally"] is True + assert result["loras"][0]["hash"] == sha256 + + +@pytest.mark.asyncio +async def test_local_cache_lora_hits_increment_base_model_counts_for_fallback(monkeypatch): + """When ALL loras hit the cache, base_model_counts is populated so the + max(counts) fallback still resolves the base model (parity with the API path).""" + provider = _RecordingProvider(raise_on_call=True) + item1 = _cache_item(name="Lora One", base_model="SDXL 1.0") + item2 = _cache_item(name="Lora Two", base_model="SDXL 1.0") + local_cache = {"hash1111111111": item1, "hash2222222222": item2} + metadata = { + "resources": [ + {"hash": "hash1111111111", "name": "Lora One", "type": "lora"}, + {"hash": "hash2222222222", "name": "Lora Two", "type": "lora"}, + ], + } + + result = await _parse_with_cache(monkeypatch, provider, metadata, local_cache=local_cache) + + assert provider.hash_calls == [] + assert len(result["loras"]) == 2 + assert result["base_model"] == "SDXL 1.0" + + +@pytest.mark.asyncio +async def test_local_cache_checkpoint_hit_does_not_increment_base_model_counts(monkeypatch): + """Parity pin: a checkpoint cache hit never contributes to base_model_counts. + Scenario 1 sets result["base_model"] directly (like the API path); scenario 2 + proves a base_model-less checkpoint added nothing to the counts fallback.""" + provider = _RecordingProvider(raise_on_call=True) + + cp_item = _cache_item(name="My Checkpoint", base_model="CP Base", model_type="Checkpoint") + lora_item = _cache_item(name="Style LoRA", base_model="Lora Base", model_type="LORA") + metadata1 = { + "resources": [ + {"hash": "cp1234567890", "name": "My Checkpoint", "type": "model"}, + {"hash": "lora123456789", "name": "Style LoRA", "type": "lora"}, + ], + } + result1 = await _parse_with_cache( + monkeypatch, + provider, + metadata1, + local_cache={"cp1234567890": cp_item, "lora123456789": lora_item}, + ) + assert result1["base_model"] == "CP Base" + + cp_no_bm = _cache_item(name="No Bm Checkpoint", base_model="", model_type="Checkpoint") + metadata2 = { + "resources": [ + {"hash": "cp9999999999", "name": "No Bm Checkpoint", "type": "model"}, + {"hash": "lora123456789", "name": "Style LoRA", "type": "lora"}, + ], + } + result2 = await _parse_with_cache( + monkeypatch, + provider, + metadata2, + local_cache={"cp9999999999": cp_no_bm, "lora123456789": lora_item}, + ) + assert result2["base_model"] == "Lora Base" + + +@pytest.mark.asyncio +async def test_local_cache_type_gate_skips_checkpoint_cache_item_in_lora_section(monkeypatch): + """A cache item whose civitai.model.type is a checkpoint is skipped in a + lora section — no entry is appended for it.""" + provider = _RecordingProvider(raise_on_call=True) + cp_item = _cache_item(name="Disguised Checkpoint", base_model="CP Base", model_type="Checkpoint") + lora_item = _cache_item(name="Real Lora", base_model="Lora Base", model_type="LORA") + metadata = { + "resources": [ + {"hash": "cp1111111111", "name": "Disguised Checkpoint", "type": "lora"}, + {"hash": "lora111111111", "name": "Real Lora", "type": "lora"}, + ], + } + + result = await _parse_with_cache( + monkeypatch, + provider, + metadata, + local_cache={"cp1111111111": cp_item, "lora111111111": lora_item}, + ) + + assert provider.hash_calls == [] + assert [l["name"] for l in result["loras"]] == ["Real Lora"] + + +@pytest.mark.asyncio +async def test_local_cache_type_gate_accepts_uppercase_lora_type(monkeypatch): + """Cache items storing the type as UPPERCASE 'LORA' must not be skipped + (false-skip guard — stored types are verbatim, VALID_LORA_TYPES is lowercase).""" + provider = _RecordingProvider(raise_on_call=True) + item = _cache_item(name="Uppercase Lora", base_model="SDXL 1.0", model_type="LORA") + local_cache = {"upcasehash12": item} + metadata = { + "resources": [ + {"hash": "upcasehash12", "name": "Uppercase Lora", "type": "lora"}, + ], + } + + result = await _parse_with_cache(monkeypatch, provider, metadata, local_cache=local_cache) + + assert provider.hash_calls == [] + assert len(result["loras"]) == 1 + assert result["loras"][0]["existsLocally"] is True + + +@pytest.mark.asyncio +async def test_local_cache_type_gate_accepts_item_without_civitai_type(monkeypatch): + """Local-only cache items without civitai type info are treated as valid.""" + provider = _RecordingProvider(raise_on_call=True) + item = { + "file_path": "/loras/local_only.safetensors", + "file_name": "Local Only", + "sha256": "00112233445566778899aabbccddeeff00112233445566778899aabbccddeeff", + "preview_url": "/previews/local_only.png", + "base_model": "SDXL 1.0", + } + local_cache = {"localonly123": item} + metadata = { + "resources": [ + {"hash": "localonly123", "name": "Local Only", "type": "lora"}, + ], + } + + result = await _parse_with_cache(monkeypatch, provider, metadata, local_cache=local_cache) + + assert provider.hash_calls == [] + assert len(result["loras"]) == 1 + assert result["loras"][0]["existsLocally"] is True + + +@pytest.mark.asyncio +async def test_local_cache_miss_calls_provider(monkeypatch): + """A hash absent from local_cache falls through to the provider as before.""" + lora_info = _make_lora_info(base_model="SDXL 1.0") + provider = _RecordingProvider(result=lora_info) + metadata = { + "resources": [ + {"hash": "missedhash123", "name": "Missed LoRA", "type": "lora", "weight": 0.8}, + ], + } + + result = await _parse_with_cache(monkeypatch, provider, metadata, local_cache={}) + + assert provider.hash_calls == ["missedhash123"] + assert len(result["loras"]) == 1 + assert result["loras"][0]["name"] == "Style LoRA" + + +@pytest.mark.asyncio +async def test_local_cache_dedup_same_hash_produces_one_entry_on_hit(monkeypatch): + """Repeated same-hash resources produce a single entry on the cache-hit path.""" + provider = _RecordingProvider(raise_on_call=True) + item = _cache_item(name="Dedup Lora", base_model="SDXL 1.0") + local_cache = {"deduphash123": item} + metadata = { + "resources": [ + {"hash": "deduphash123", "name": "Dedup Lora", "type": "lora"}, + {"hash": "deduphash123", "name": "Dedup Lora", "type": "lora"}, + ], + } + + result = await _parse_with_cache(monkeypatch, provider, metadata, local_cache=local_cache) + + assert provider.hash_calls == [] + assert len(result["loras"]) == 1 + + +@pytest.mark.asyncio +async def test_local_cache_dedup_same_hash_produces_one_entry_on_miss(monkeypatch): + """Repeated same-hash resources produce a single entry on the cache-miss path.""" + provider = _RecordingProvider(result=_make_lora_info(base_model="SDXL 1.0")) + metadata = { + "resources": [ + {"hash": "missdedup123", "name": "Dedup Miss", "type": "lora"}, + {"hash": "missdedup123", "name": "Dedup Miss", "type": "lora"}, + ], + } + + result = await _parse_with_cache(monkeypatch, provider, metadata, local_cache={}) + + assert provider.hash_calls == ["missdedup123"] + assert len(result["loras"]) == 1 +