mirror of
https://github.com/willmiao/ComfyUI-Lora-Manager.git
synced 2026-10-02 16:25:33 -03:00
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.
This commit is contained in:
@@ -29,6 +29,7 @@ class ImportExampleImagesUseCase:
|
|||||||
|
|
||||||
async def execute(self, request: web.Request) -> Dict[str, Any]:
|
async def execute(self, request: web.Request) -> Dict[str, Any]:
|
||||||
model_hash: str | None = None
|
model_hash: str | None = None
|
||||||
|
model_path: str | None = None
|
||||||
files_to_import: List[str] = []
|
files_to_import: List[str] = []
|
||||||
temp_files: 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
|
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":
|
if first_field and first_field.name == "model_hash":
|
||||||
model_hash = await first_field.text()
|
model_hash = await first_field.text()
|
||||||
|
elif first_field and first_field.name == "model_path":
|
||||||
|
model_path = await first_field.text()
|
||||||
else:
|
else:
|
||||||
# Support clients that send files first and hash later
|
# Support clients that send files first and hash later
|
||||||
if first_field is not None:
|
if first_field is not None:
|
||||||
@@ -49,13 +52,22 @@ class ImportExampleImagesUseCase:
|
|||||||
field = cast(BodyPartReader, raw_field)
|
field = cast(BodyPartReader, raw_field)
|
||||||
if field.name == "model_hash" and not model_hash:
|
if field.name == "model_hash" and not model_hash:
|
||||||
model_hash = await field.text()
|
model_hash = await field.text()
|
||||||
|
elif field.name == "model_path" and not model_path:
|
||||||
|
model_path = await field.text()
|
||||||
elif field.name == "files":
|
elif field.name == "files":
|
||||||
await self._collect_upload_file(field, files_to_import, temp_files)
|
await self._collect_upload_file(field, files_to_import, temp_files)
|
||||||
else:
|
else:
|
||||||
data = await request.json()
|
data = await request.json()
|
||||||
model_hash = data.get("model_hash")
|
model_hash = data.get("model_hash")
|
||||||
|
model_path = data.get("model_path")
|
||||||
files_to_import = list(data.get("file_paths", []))
|
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:
|
if not model_hash:
|
||||||
raise ImportExampleImagesValidationError("Missing model_hash parameter")
|
raise ImportExampleImagesValidationError("Missing model_hash parameter")
|
||||||
result = await self._processor.import_images(model_hash, files_to_import)
|
result = await self._processor.import_images(model_hash, files_to_import)
|
||||||
|
|||||||
@@ -26,6 +26,42 @@ class ExampleImagesValidationError(ExampleImagesImportError):
|
|||||||
class ExampleImagesProcessor:
|
class ExampleImagesProcessor:
|
||||||
"""Processes and manipulates example images"""
|
"""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
|
@staticmethod
|
||||||
def generate_short_id(length=8):
|
def generate_short_id(length=8):
|
||||||
"""Generate a short random alphanumeric identifier"""
|
"""Generate a short random alphanumeric identifier"""
|
||||||
@@ -450,15 +486,11 @@ class ExampleImagesProcessor:
|
|||||||
raise ExampleImagesValidationError('No example images path configured')
|
raise ExampleImagesValidationError('No example images path configured')
|
||||||
|
|
||||||
# Find the model and get current metadata
|
# 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
|
model_data = None
|
||||||
scanner = None
|
scanner = None
|
||||||
|
|
||||||
# Check both scanners to find the model
|
# Check every scanner to find the model
|
||||||
for scan_obj in [lora_scanner, checkpoint_scanner, embedding_scanner]:
|
for scan_obj in await ExampleImagesProcessor._model_scanners():
|
||||||
cache = await scan_obj.get_cached_data()
|
cache = await scan_obj.get_cached_data()
|
||||||
for item in cache.raw_data:
|
for item in cache.raw_data:
|
||||||
if item.get('sha256') == model_hash:
|
if item.get('sha256') == model_hash:
|
||||||
@@ -536,6 +568,7 @@ class ExampleImagesProcessor:
|
|||||||
'errors': errors,
|
'errors': errors,
|
||||||
'regular_images': regular_images,
|
'regular_images': regular_images,
|
||||||
'custom_images': custom_images,
|
'custom_images': custom_images,
|
||||||
|
'model_hash': model_hash,
|
||||||
"model_file_path": model_data.get('file_path', ''),
|
"model_file_path": model_data.get('file_path', ''),
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -577,15 +610,11 @@ class ExampleImagesProcessor:
|
|||||||
}, status=400)
|
}, status=400)
|
||||||
|
|
||||||
# Find the model and get current metadata
|
# 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
|
model_data = None
|
||||||
scanner = None
|
scanner = None
|
||||||
|
|
||||||
# Check both scanners to find the model
|
# Check every scanner to find the model
|
||||||
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):
|
if scan_obj.has_hash(model_hash):
|
||||||
cache = await scan_obj.get_cached_data()
|
cache = await scan_obj.get_cached_data()
|
||||||
for item in cache.raw_data:
|
for item in cache.raw_data:
|
||||||
@@ -737,14 +766,10 @@ class ExampleImagesProcessor:
|
|||||||
)
|
)
|
||||||
|
|
||||||
try:
|
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
|
model_data = None
|
||||||
scanner = 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):
|
if scan_obj.has_hash(model_hash):
|
||||||
cache = await scan_obj.get_cached_data()
|
cache = await scan_obj.get_cached_data()
|
||||||
for item in cache.raw_data:
|
for item in cache.raw_data:
|
||||||
|
|||||||
@@ -1162,10 +1162,20 @@ async function handleImportFiles(files, modelHash, importContainer) {
|
|||||||
let successCount = 0;
|
let successCount = 0;
|
||||||
const errors = [];
|
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) {
|
for (const file of validFiles) {
|
||||||
try {
|
try {
|
||||||
const formData = new FormData();
|
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);
|
formData.append('files', file);
|
||||||
|
|
||||||
const response = await fetch('/api/lm/import-example-images', {
|
const response = await fetch('/api/lm/import-example-images', {
|
||||||
@@ -1192,8 +1202,17 @@ async function handleImportFiles(files, modelHash, importContainer) {
|
|||||||
|
|
||||||
const result = lastSuccessResult;
|
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
|
// 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();
|
const updatedFilesResult = await updatedFilesResponse.json();
|
||||||
|
|
||||||
if (!updatedFilesResult.success) {
|
if (!updatedFilesResult.success) {
|
||||||
@@ -1221,7 +1240,7 @@ async function handleImportFiles(files, modelHash, importContainer) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
// Initialize the import UI for the new content
|
// Initialize the import UI for the new content
|
||||||
initExampleImport(modelHash, showcaseTab);
|
initExampleImport(effectiveHash, showcaseTab);
|
||||||
|
|
||||||
if (errors.length > 0) {
|
if (errors.length > 0) {
|
||||||
showToast('toast.import.imagesPartial', { success: successCount, failed: errors.length }, 'warning');
|
showToast('toast.import.imagesPartial', { success: successCount, failed: errors.length }, 'warning');
|
||||||
|
|||||||
@@ -146,6 +146,8 @@ class StubExampleImagesProcessor(ExampleImagesProcessor):
|
|||||||
self.calls: List[Dict[str, Any]] = []
|
self.calls: List[Dict[str, Any]] = []
|
||||||
self.error: Optional[str] = None
|
self.error: Optional[str] = None
|
||||||
self.response: Dict[str, Any] = {"success": True}
|
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]
|
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})
|
self.calls.append({"model_hash": model_hash, "files": files})
|
||||||
@@ -155,6 +157,10 @@ class StubExampleImagesProcessor(ExampleImagesProcessor):
|
|||||||
raise ExampleImagesImportError("boom")
|
raise ExampleImagesImportError("boom")
|
||||||
return self.response
|
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:
|
async def test_auto_organize_use_case_executes_with_lock() -> None:
|
||||||
file_service = StubFileService()
|
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]
|
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:
|
class StubLifecycleService:
|
||||||
def __init__(self, scanner: Optional[MockScanner] = None) -> None:
|
def __init__(self, scanner: Optional[MockScanner] = None) -> None:
|
||||||
self.renames: List[Dict[str, str]] = []
|
self.renames: List[Dict[str, str]] = []
|
||||||
|
|||||||
@@ -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_lora_scanner", classmethod(_return_scanner))
|
||||||
monkeypatch.setattr(processor_module.ServiceRegistry, "get_checkpoint_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_embedding_scanner", classmethod(_return_scanner))
|
||||||
|
monkeypatch.setattr(processor_module.ServiceRegistry, "get_other_scanner", classmethod(_return_scanner))
|
||||||
|
|
||||||
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_lora_scanner", classmethod(_empty_scanner))
|
||||||
monkeypatch.setattr(processor_module.ServiceRegistry, "get_checkpoint_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_embedding_scanner", classmethod(_empty_scanner))
|
||||||
|
monkeypatch.setattr(processor_module.ServiceRegistry, "get_other_scanner", classmethod(_empty_scanner))
|
||||||
|
|
||||||
with pytest.raises(processor_module.ExampleImagesImportError):
|
with pytest.raises(processor_module.ExampleImagesImportError):
|
||||||
await processor_module.ExampleImagesProcessor.import_images("a" * 64, [str(tmp_path / "missing.png")])
|
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_lora_scanner", classmethod(_return_scanner))
|
||||||
monkeypatch.setattr(processor_module.ServiceRegistry, "get_checkpoint_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_embedding_scanner", classmethod(_return_scanner))
|
||||||
|
monkeypatch.setattr(processor_module.ServiceRegistry, "get_other_scanner", classmethod(_return_scanner))
|
||||||
|
|
||||||
model_folder = get_model_folder(model_hash)
|
model_folder = get_model_folder(model_hash)
|
||||||
os.makedirs(model_folder, exist_ok=True)
|
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]
|
_, _, updated_metadata = scanner.updated[-1]
|
||||||
assert updated_metadata["civitai"]["images"] == existing_metadata["civitai"]["images"]
|
assert updated_metadata["civitai"]["images"] == existing_metadata["civitai"]["images"]
|
||||||
assert updated_metadata["civitai"]["customImages"] == []
|
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"
|
||||||
|
)
|
||||||
|
== ""
|
||||||
|
)
|
||||||
|
|||||||
Reference in New Issue
Block a user