From 2193ec8f38de7f5bdbffc80cb34de250b9d01705 Mon Sep 17 00:00:00 2001 From: Will Miao Date: Thu, 1 Oct 2026 10:16:34 +0800 Subject: [PATCH] fix(example-images): support importing for Other models and pending hashes Two gaps kept 'import example images' from working on the Other page: - import_images/delete_custom_image/set_example_image_nsfw_level only searched the lora/checkpoint/embedding scanners, so Other-category models were never found ('Model with hash ... not found in cache'). All three now go through a shared scanner list that includes the Other scanner. - Other models (and fresh checkpoints) carry hash_status=pending with an empty sha256, so the frontend sent an empty model_hash and the import was rejected with 'Missing model_hash parameter'. The modal now also sends the model's file path, and the import use case resolves the hash on demand via the scanner's lazy-hash calculation. The resolved hash is returned to the UI and persisted on the showcase element so follow-up operations target the same hash-keyed folder. --- .../import_example_images_use_case.py | 12 ++ py/utils/example_images_processor.py | 65 ++++++--- .../shared/showcase/ShowcaseView.js | 25 +++- tests/services/test_use_cases.py | 34 +++++ .../test_example_images_processor_unit.py | 135 ++++++++++++++++++ 5 files changed, 248 insertions(+), 23 deletions(-) diff --git a/py/services/use_cases/example_images/import_example_images_use_case.py b/py/services/use_cases/example_images/import_example_images_use_case.py index e7d51614..001a8ef4 100644 --- a/py/services/use_cases/example_images/import_example_images_use_case.py +++ b/py/services/use_cases/example_images/import_example_images_use_case.py @@ -29,6 +29,7 @@ class ImportExampleImagesUseCase: async def execute(self, request: web.Request) -> Dict[str, Any]: model_hash: str | None = None + model_path: str | None = None files_to_import: List[str] = [] temp_files: List[str] = [] @@ -40,6 +41,8 @@ class ImportExampleImagesUseCase: first_field = cast(BodyPartReader, first_field_raw) if first_field_raw is not None else None if first_field and first_field.name == "model_hash": model_hash = await first_field.text() + elif first_field and first_field.name == "model_path": + model_path = await first_field.text() else: # Support clients that send files first and hash later if first_field is not None: @@ -49,13 +52,22 @@ class ImportExampleImagesUseCase: field = cast(BodyPartReader, raw_field) if field.name == "model_hash" and not model_hash: model_hash = await field.text() + elif field.name == "model_path" and not model_path: + model_path = await field.text() elif field.name == "files": await self._collect_upload_file(field, files_to_import, temp_files) else: data = await request.json() model_hash = data.get("model_hash") + model_path = data.get("model_path") files_to_import = list(data.get("file_paths", [])) + # Models with a deferred hash (checkpoints, Other) send an empty + # model_hash; locate them by file path and compute the hash on + # demand, since example-image folders are keyed by hash. + if not model_hash and model_path: + model_hash = await self._processor.resolve_hash_for_file_path(model_path) + if not model_hash: raise ImportExampleImagesValidationError("Missing model_hash parameter") result = await self._processor.import_images(model_hash, files_to_import) diff --git a/py/utils/example_images_processor.py b/py/utils/example_images_processor.py index ed6fed54..e06c2c31 100644 --- a/py/utils/example_images_processor.py +++ b/py/utils/example_images_processor.py @@ -26,6 +26,42 @@ class ExampleImagesValidationError(ExampleImagesImportError): class ExampleImagesProcessor: """Processes and manipulates example images""" + @staticmethod + async def _model_scanners() -> list: + """Return every scanner whose models can carry example images.""" + return [ + await ServiceRegistry.get_lora_scanner(), + await ServiceRegistry.get_checkpoint_scanner(), + await ServiceRegistry.get_embedding_scanner(), + await ServiceRegistry.get_other_scanner(), + ] + + @staticmethod + async def resolve_hash_for_file_path(file_path: str) -> str: + """Return the SHA256 for a cached model file, computing it on demand. + + Checkpoint and Other scanners record ``hash_status="pending"`` with an + empty sha256 until something needs the hash; importing example images + is such a moment because example folders are keyed by hash. Returns + ``''`` when the file is unknown to every scanner or hashing failed. + """ + if not file_path: + return '' + normalized = file_path.replace(os.sep, '/') + for scanner in await ExampleImagesProcessor._model_scanners(): + cache = await scanner.get_cached_data() + for item in cache.raw_data: + if item.get('file_path') != normalized: + continue + sha256 = (item.get('sha256') or '').strip() + if sha256 and item.get('hash_status', 'completed') == 'completed': + return sha256 + calculate = getattr(scanner, 'calculate_hash_for_model', None) + if calculate is None: + return sha256 + return (await calculate(normalized)) or '' + return '' + @staticmethod def generate_short_id(length=8): """Generate a short random alphanumeric identifier""" @@ -450,15 +486,11 @@ class ExampleImagesProcessor: raise ExampleImagesValidationError('No example images path configured') # Find the model and get current metadata - lora_scanner = await ServiceRegistry.get_lora_scanner() - checkpoint_scanner = await ServiceRegistry.get_checkpoint_scanner() - embedding_scanner = await ServiceRegistry.get_embedding_scanner() - model_data = None scanner = None - # Check both scanners to find the model - for scan_obj in [lora_scanner, checkpoint_scanner, embedding_scanner]: + # Check every scanner to find the model + for scan_obj in await ExampleImagesProcessor._model_scanners(): cache = await scan_obj.get_cached_data() for item in cache.raw_data: if item.get('sha256') == model_hash: @@ -536,6 +568,7 @@ class ExampleImagesProcessor: 'errors': errors, 'regular_images': regular_images, 'custom_images': custom_images, + 'model_hash': model_hash, "model_file_path": model_data.get('file_path', ''), } @@ -577,15 +610,11 @@ class ExampleImagesProcessor: }, status=400) # Find the model and get current metadata - lora_scanner = await ServiceRegistry.get_lora_scanner() - checkpoint_scanner = await ServiceRegistry.get_checkpoint_scanner() - embedding_scanner = await ServiceRegistry.get_embedding_scanner() - model_data = None scanner = None - - # Check both scanners to find the model - for scan_obj in [lora_scanner, checkpoint_scanner, embedding_scanner]: + + # Check every scanner to find the model + for scan_obj in await ExampleImagesProcessor._model_scanners(): if scan_obj.has_hash(model_hash): cache = await scan_obj.get_cached_data() for item in cache.raw_data: @@ -595,13 +624,13 @@ class ExampleImagesProcessor: break if model_data: break - + if not model_data: return web.json_response({ 'success': False, 'error': f"Model with hash {model_hash} not found in cache" }, status=404) - + await MetadataManager.hydrate_model_data(model_data) civitai_data = model_data.setdefault('civitai', {}) custom_images = civitai_data.get('customImages') @@ -737,14 +766,10 @@ class ExampleImagesProcessor: ) try: - lora_scanner = await ServiceRegistry.get_lora_scanner() - checkpoint_scanner = await ServiceRegistry.get_checkpoint_scanner() - embedding_scanner = await ServiceRegistry.get_embedding_scanner() - model_data = None scanner = None - for scan_obj in [lora_scanner, checkpoint_scanner, embedding_scanner]: + for scan_obj in await ExampleImagesProcessor._model_scanners(): if scan_obj.has_hash(model_hash): cache = await scan_obj.get_cached_data() for item in cache.raw_data: diff --git a/static/js/components/shared/showcase/ShowcaseView.js b/static/js/components/shared/showcase/ShowcaseView.js index d9baf6e6..d5476513 100644 --- a/static/js/components/shared/showcase/ShowcaseView.js +++ b/static/js/components/shared/showcase/ShowcaseView.js @@ -1162,10 +1162,20 @@ async function handleImportFiles(files, modelHash, importContainer) { let successCount = 0; const errors = []; + // The showcase section carries the freshest hash (an import may have + // resolved a deferred hash) plus the file path the backend needs to + // locate models whose hash is still pending. + const showcaseSection = document.querySelector('.showcase-section'); + const modelPath = showcaseSection?.dataset.filepath || ''; + const currentHash = showcaseSection?.dataset.modelHash || modelHash; + for (const file of validFiles) { try { const formData = new FormData(); - formData.append('model_hash', modelHash); + formData.append('model_hash', currentHash); + if (modelPath) { + formData.append('model_path', modelPath); + } formData.append('files', file); const response = await fetch('/api/lm/import-example-images', { @@ -1192,8 +1202,17 @@ async function handleImportFiles(files, modelHash, importContainer) { const result = lastSuccessResult; + // A model with a deferred hash (checkpoint / Other) is imported via + // model_path; the backend resolves and returns the real hash. Persist + // it so every follow-up (file list, NSFW toggle, delete, re-render) + // targets the hash-keyed example folder. + const effectiveHash = result.model_hash || currentHash; + if (showcaseSection && effectiveHash) { + showcaseSection.dataset.modelHash = effectiveHash; + } + // Get updated local files - const updatedFilesResponse = await fetch(`/api/lm/example-image-files?model_hash=${modelHash}`); + const updatedFilesResponse = await fetch(`/api/lm/example-image-files?model_hash=${effectiveHash}`); const updatedFilesResult = await updatedFilesResponse.json(); if (!updatedFilesResult.success) { @@ -1221,7 +1240,7 @@ async function handleImportFiles(files, modelHash, importContainer) { } // Initialize the import UI for the new content - initExampleImport(modelHash, showcaseTab); + initExampleImport(effectiveHash, showcaseTab); if (errors.length > 0) { showToast('toast.import.imagesPartial', { success: successCount, failed: errors.length }, 'warning'); diff --git a/tests/services/test_use_cases.py b/tests/services/test_use_cases.py index 7d537cc7..e24d4817 100644 --- a/tests/services/test_use_cases.py +++ b/tests/services/test_use_cases.py @@ -146,6 +146,8 @@ class StubExampleImagesProcessor(ExampleImagesProcessor): self.calls: List[Dict[str, Any]] = [] self.error: Optional[str] = None self.response: Dict[str, Any] = {"success": True} + self.resolved_hash: Optional[str] = None + self.resolve_calls: List[str] = [] async def import_images(self, model_hash: str, files: List[str]) -> Dict[str, Any]: # pyright: ignore[reportIncompatibleMethodOverride] self.calls.append({"model_hash": model_hash, "files": files}) @@ -155,6 +157,10 @@ class StubExampleImagesProcessor(ExampleImagesProcessor): raise ExampleImagesImportError("boom") return self.response + async def resolve_hash_for_file_path(self, file_path: str) -> str: + self.resolve_calls.append(file_path) + return self.resolved_hash or "" + async def test_auto_organize_use_case_executes_with_lock() -> None: file_service = StubFileService() @@ -506,6 +512,34 @@ async def test_import_example_images_use_case_propagates_generic_error() -> None await use_case.execute(request) # pyright: ignore[reportArgumentType] +async def test_import_example_images_use_case_resolves_hash_from_model_path() -> None: + """Models with a deferred hash (checkpoints, Other) send model_path + instead of model_hash; the use case must resolve the hash on demand.""" + processor = StubExampleImagesProcessor() + processor.resolved_hash = "f" * 64 + use_case = ImportExampleImagesUseCase(processor=processor) + + request = DummyJsonRequest( + {"model_path": "/models/vae/x.safetensors", "file_paths": ["/tmp/file"]} + ) + result = await use_case.execute(request) # pyright: ignore[reportArgumentType] + + assert processor.resolve_calls == ["/models/vae/x.safetensors"] + assert processor.calls == [{"model_hash": "f" * 64, "files": ["/tmp/file"]}] + assert result == {"success": True} + + +async def test_import_example_images_use_case_rejects_unresolvable_model_path() -> None: + processor = StubExampleImagesProcessor() + use_case = ImportExampleImagesUseCase(processor=processor) + request = DummyJsonRequest( + {"model_path": "/models/unknown.safetensors", "file_paths": []} + ) + + with pytest.raises(ImportExampleImagesValidationError): + await use_case.execute(request) # pyright: ignore[reportArgumentType] + + class StubLifecycleService: def __init__(self, scanner: Optional[MockScanner] = None) -> None: self.renames: List[Dict[str, str]] = [] diff --git a/tests/utils/test_example_images_processor_unit.py b/tests/utils/test_example_images_processor_unit.py index 229d5072..02da9d21 100644 --- a/tests/utils/test_example_images_processor_unit.py +++ b/tests/utils/test_example_images_processor_unit.py @@ -183,6 +183,7 @@ def stub_scanners(monkeypatch: pytest.MonkeyPatch, tmp_path) -> StubScanner: monkeypatch.setattr(processor_module.ServiceRegistry, "get_lora_scanner", classmethod(_return_scanner)) monkeypatch.setattr(processor_module.ServiceRegistry, "get_checkpoint_scanner", classmethod(_return_scanner)) monkeypatch.setattr(processor_module.ServiceRegistry, "get_embedding_scanner", classmethod(_return_scanner)) + monkeypatch.setattr(processor_module.ServiceRegistry, "get_other_scanner", classmethod(_return_scanner)) return scanner @@ -240,6 +241,7 @@ async def test_import_images_raises_when_model_not_found(monkeypatch: pytest.Mon monkeypatch.setattr(processor_module.ServiceRegistry, "get_lora_scanner", classmethod(_empty_scanner)) monkeypatch.setattr(processor_module.ServiceRegistry, "get_checkpoint_scanner", classmethod(_empty_scanner)) monkeypatch.setattr(processor_module.ServiceRegistry, "get_embedding_scanner", classmethod(_empty_scanner)) + monkeypatch.setattr(processor_module.ServiceRegistry, "get_other_scanner", classmethod(_empty_scanner)) with pytest.raises(processor_module.ExampleImagesImportError): await processor_module.ExampleImagesProcessor.import_images("a" * 64, [str(tmp_path / "missing.png")]) @@ -288,6 +290,7 @@ async def test_delete_custom_image_preserves_existing_metadata(monkeypatch: pyte monkeypatch.setattr(processor_module.ServiceRegistry, "get_lora_scanner", classmethod(_return_scanner)) monkeypatch.setattr(processor_module.ServiceRegistry, "get_checkpoint_scanner", classmethod(_return_scanner)) monkeypatch.setattr(processor_module.ServiceRegistry, "get_embedding_scanner", classmethod(_return_scanner)) + monkeypatch.setattr(processor_module.ServiceRegistry, "get_other_scanner", classmethod(_return_scanner)) model_folder = get_model_folder(model_hash) os.makedirs(model_folder, exist_ok=True) @@ -329,3 +332,135 @@ async def test_delete_custom_image_preserves_existing_metadata(monkeypatch: pyte _, _, updated_metadata = scanner.updated[-1] assert updated_metadata["civitai"]["images"] == existing_metadata["civitai"]["images"] assert updated_metadata["civitai"]["customImages"] == [] + + +def _patch_scanner_getters(monkeypatch: pytest.MonkeyPatch, **scanners) -> None: + """Point every model-scanner getter at the given stub (empty by default).""" + + async def _empty(cls=None): + return StubScanner([]) + + for name in ( + "get_lora_scanner", + "get_checkpoint_scanner", + "get_embedding_scanner", + "get_other_scanner", + ): + stub = scanners.get(name) + if stub is None: + getter = _empty + else: + async def getter(cls=None, _stub=stub): # noqa: B023 - bound per iteration + return _stub + monkeypatch.setattr( + processor_module.ServiceRegistry, name, classmethod(getter) + ) + + +@pytest.mark.asyncio +async def test_import_images_finds_model_in_other_scanner( + monkeypatch: pytest.MonkeyPatch, tmp_path +) -> None: + """Other-category models (VAEs, text encoders) live in the Other scanner; + importing example images must search it too.""" + settings_manager = get_settings_manager() + settings_manager.settings["example_images_path"] = str(tmp_path / "examples") + settings_manager.settings["libraries"] = {"default": {}} + settings_manager.settings["active_library"] = "default" + + model_hash = "b" * 64 + model_data = { + "sha256": model_hash, + "model_name": "VAE", + "file_path": str(tmp_path / "vae.safetensors"), + "civitai": {}, + } + other_scanner = StubScanner([model_data]) + _patch_scanner_getters(monkeypatch, get_other_scanner=other_scanner) + + source_file = tmp_path / "upload.png" + source_file.write_bytes(b"PNG data") + monkeypatch.setattr( + processor_module.ExampleImagesProcessor, + "generate_short_id", + staticmethod(lambda: "short"), + ) + + recorded: Dict[str, Any] = {} + + async def fake_update_metadata(model_hash, model_data, scanner, paths): + recorded["scanner"] = scanner + return [], [] + + monkeypatch.setattr( + processor_module.MetadataUpdater, + "update_metadata_after_import", + staticmethod(fake_update_metadata), + ) + + result = await processor_module.ExampleImagesProcessor.import_images( + model_hash, [str(source_file)] + ) + + assert result["success"] is True + assert result["model_hash"] == model_hash + assert recorded["scanner"] is other_scanner + + +@pytest.mark.asyncio +async def test_resolve_hash_for_file_path_computes_pending_hash( + monkeypatch: pytest.MonkeyPatch, tmp_path +) -> None: + """A model whose hash is still pending gets it computed on demand.""" + model_path = str(tmp_path / "encoder.safetensors").replace(os.sep, "/") + item = {"file_path": model_path, "sha256": "", "hash_status": "pending"} + + calculated: list[str] = [] + + class PendingScanner(StubScanner): + async def calculate_hash_for_model(self, file_path: str): + calculated.append(file_path) + return "f" * 64 + + _patch_scanner_getters(monkeypatch, get_other_scanner=PendingScanner([item])) + + result = await processor_module.ExampleImagesProcessor.resolve_hash_for_file_path( + model_path + ) + + assert result == "f" * 64 + assert calculated == [model_path] + + +@pytest.mark.asyncio +async def test_resolve_hash_for_file_path_returns_completed_hash_without_recompute( + monkeypatch: pytest.MonkeyPatch, tmp_path +) -> None: + model_path = str(tmp_path / "model.safetensors").replace(os.sep, "/") + item = {"file_path": model_path, "sha256": "c" * 64, "hash_status": "completed"} + + class EagerScanner(StubScanner): + async def calculate_hash_for_model(self, file_path: str): # pragma: no cover + raise AssertionError("must not recompute a completed hash") + + _patch_scanner_getters(monkeypatch, get_lora_scanner=EagerScanner([item])) + + result = await processor_module.ExampleImagesProcessor.resolve_hash_for_file_path( + model_path + ) + + assert result == "c" * 64 + + +@pytest.mark.asyncio +async def test_resolve_hash_for_file_path_unknown_file_returns_empty( + monkeypatch: pytest.MonkeyPatch, +) -> None: + _patch_scanner_getters(monkeypatch) + + assert ( + await processor_module.ExampleImagesProcessor.resolve_hash_for_file_path( + "/models/nope.safetensors" + ) + == "" + )