mirror of
https://github.com/willmiao/ComfyUI-Lora-Manager.git
synced 2026-09-22 03:24:09 -03:00
Compare commits
41
Commits
v1.2.0
...
4bf9a4b640
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
4bf9a4b640 | ||
|
|
c5088772e8 | ||
|
|
56acefbd6c | ||
|
|
5ab06c4aae | ||
|
|
c11f4b5c68 | ||
|
|
86376284f4 | ||
|
|
2b8a2fc7d8 | ||
|
|
f26e1b41c8 | ||
|
|
c1671af99f | ||
|
|
ac7707d0f6 | ||
|
|
381cd710a2 | ||
|
|
ad0d18cb79 | ||
|
|
7980ee77d0 | ||
|
|
916b8bb327 | ||
|
|
87e3d4dea9 | ||
|
|
76a913f5e0 | ||
|
|
d8c192e647 | ||
|
|
c453437620 | ||
|
|
720fa6d909 | ||
|
|
b4f71089f4 | ||
|
|
83e6657ead | ||
|
|
7ea6df4111 | ||
|
|
d9ab92602a | ||
|
|
5ffadaed31 | ||
|
|
24f5f7df5d | ||
|
|
daf01fb1d6 | ||
|
|
0f11b6def9 | ||
|
|
7df83f44b8 | ||
|
|
169fa7bed6 | ||
|
|
027b504fe8 | ||
|
|
186ef4da78 | ||
|
|
dc674098e7 | ||
|
|
9087b4b07c | ||
|
|
8e45c22d7a | ||
|
|
191c4e03cd | ||
|
|
ab4154c57d | ||
|
|
28e93d12ff | ||
|
|
75e63c758b | ||
|
|
823f71f269 | ||
|
|
042dd4088d | ||
|
|
eaa791a9eb |
@@ -25,6 +25,7 @@ model_cache/
|
|||||||
reasonix.toml
|
reasonix.toml
|
||||||
.reasonix/
|
.reasonix/
|
||||||
.codegraph/
|
.codegraph/
|
||||||
|
.playwright-mcp/
|
||||||
|
|
||||||
# Vue widgets development cache (but keep build output)
|
# Vue widgets development cache (but keep build output)
|
||||||
vue-widgets/node_modules/
|
vue-widgets/node_modules/
|
||||||
|
|||||||
@@ -31,7 +31,7 @@ COVERAGE_FILE=coverage/backend/.coverage pytest \
|
|||||||
--cov-report=xml:coverage/backend/coverage.xml
|
--cov-report=xml:coverage/backend/coverage.xml
|
||||||
```
|
```
|
||||||
|
|
||||||
### Frontend Development (Standalone Web UI)
|
### Frontend Development (LoRA Manager Web UI)
|
||||||
|
|
||||||
```bash
|
```bash
|
||||||
npm install
|
npm install
|
||||||
@@ -154,9 +154,9 @@ npm run test:coverage # Generate coverage report
|
|||||||
|
|
||||||
## Frontend UI Architecture
|
## Frontend UI Architecture
|
||||||
|
|
||||||
### 1. Standalone Web UI
|
### 1. LoRA Manager Web UI
|
||||||
- Location: `./static/` and `./templates/`
|
- Location: `./static/` and `./templates/`
|
||||||
- Tech: Vanilla JS + CSS, served by standalone server
|
- Tech: Vanilla JS + CSS, served by the hosting server (ComfyUI app in plugin mode, `standalone.py` in standalone mode)
|
||||||
- Tests via npm in root directory
|
- Tests via npm in root directory
|
||||||
|
|
||||||
### 2. ComfyUI Custom Node Widgets
|
### 2. ComfyUI Custom Node Widgets
|
||||||
|
|||||||
+2249
-2222
File diff suppressed because it is too large
Load Diff
+29
-2
@@ -449,6 +449,12 @@
|
|||||||
"compact": "7 (1080p), 8 (2K), 10 (4K)"
|
"compact": "7 (1080p), 8 (2K), 10 (4K)"
|
||||||
},
|
},
|
||||||
"displayDensityWarning": "Warning: Higher densities may cause performance issues on systems with limited resources.",
|
"displayDensityWarning": "Warning: Higher densities may cause performance issues on systems with limited resources.",
|
||||||
|
"recipesLayout": "Recipes Layout",
|
||||||
|
"recipesLayoutHelp": "Choose how recipe cards are arranged: a uniform grid or a masonry (Pinterest-style) layout that preserves each image's aspect ratio.",
|
||||||
|
"recipesLayoutOptions": {
|
||||||
|
"grid": "Grid",
|
||||||
|
"masonry": "Masonry"
|
||||||
|
},
|
||||||
"showFolderSidebar": "Show Folder Sidebar",
|
"showFolderSidebar": "Show Folder Sidebar",
|
||||||
"showFolderSidebarHelp": "Toggle the folder navigation sidebar on model pages. When disabled, the sidebar and hover area stay hidden.",
|
"showFolderSidebarHelp": "Toggle the folder navigation sidebar on model pages. When disabled, the sidebar and hover area stay hidden.",
|
||||||
"cardInfoDisplay": "Card Info Display",
|
"cardInfoDisplay": "Card Info Display",
|
||||||
@@ -678,6 +684,7 @@
|
|||||||
"deepseek": "DeepSeek",
|
"deepseek": "DeepSeek",
|
||||||
"groq": "Groq",
|
"groq": "Groq",
|
||||||
"openrouter": "OpenRouter",
|
"openrouter": "OpenRouter",
|
||||||
|
"google": "Gemini",
|
||||||
"opencode-go": "OpenCode Go",
|
"opencode-go": "OpenCode Go",
|
||||||
"custom": "Custom (OpenAI-compatible)"
|
"custom": "Custom (OpenAI-compatible)"
|
||||||
},
|
},
|
||||||
@@ -714,7 +721,9 @@
|
|||||||
"versionsCount": "Local Versions",
|
"versionsCount": "Local Versions",
|
||||||
"versionsCountDesc": "Most versions first",
|
"versionsCountDesc": "Most versions first",
|
||||||
"versionsCountAsc": "Fewest versions first",
|
"versionsCountAsc": "Fewest versions first",
|
||||||
"versionIdDesc": "Newest version first"
|
"versionIdDesc": "Newest version first",
|
||||||
|
"random": "Random",
|
||||||
|
"randomAction": "Randomize (shuffle)"
|
||||||
},
|
},
|
||||||
"refresh": {
|
"refresh": {
|
||||||
"title": "Refresh model list",
|
"title": "Refresh model list",
|
||||||
@@ -771,6 +780,8 @@
|
|||||||
"deleteAll": "Delete Selected",
|
"deleteAll": "Delete Selected",
|
||||||
"downloadMissingLoras": "Download Missing LoRAs",
|
"downloadMissingLoras": "Download Missing LoRAs",
|
||||||
"downloadExamples": "Download Example Images",
|
"downloadExamples": "Download Example Images",
|
||||||
|
"downloadMissingExamples": "Download Missing",
|
||||||
|
"reprocessExamples": "Re-process All",
|
||||||
"clear": "Clear Selection",
|
"clear": "Clear Selection",
|
||||||
"skipMetadataRefreshCount": "Skip ({count} models)",
|
"skipMetadataRefreshCount": "Skip ({count} models)",
|
||||||
"resumeMetadataRefreshCount": "Resume ({count} models)",
|
"resumeMetadataRefreshCount": "Resume ({count} models)",
|
||||||
@@ -806,6 +817,8 @@
|
|||||||
"sendToWorkflowReplace": "Send to Workflow (Replace)",
|
"sendToWorkflowReplace": "Send to Workflow (Replace)",
|
||||||
"openExamples": "Open Examples Folder",
|
"openExamples": "Open Examples Folder",
|
||||||
"downloadExamples": "Download Example Images",
|
"downloadExamples": "Download Example Images",
|
||||||
|
"downloadMissingExamples": "Download Missing",
|
||||||
|
"reprocessExamples": "Re-process All",
|
||||||
"replacePreview": "Replace Preview",
|
"replacePreview": "Replace Preview",
|
||||||
"setContentRating": "Set Content Rating",
|
"setContentRating": "Set Content Rating",
|
||||||
"moveToFolder": "Move to Folder",
|
"moveToFolder": "Move to Folder",
|
||||||
@@ -1575,6 +1588,21 @@
|
|||||||
"downloadCsv": "Download CSV",
|
"downloadCsv": "Download CSV",
|
||||||
"columnModelName": "Model Name",
|
"columnModelName": "Model Name",
|
||||||
"columnError": "Error"
|
"columnError": "Error"
|
||||||
|
},
|
||||||
|
"downloadBatchSummary": {
|
||||||
|
"title": "Batch Download Summary",
|
||||||
|
"statSuccess": "Success",
|
||||||
|
"statFailed": "Failed",
|
||||||
|
"statTotal": "Total",
|
||||||
|
"successMessage": "All {count} models downloaded successfully",
|
||||||
|
"completedWithErrors": "Completed with errors",
|
||||||
|
"failed": "Download failed",
|
||||||
|
"failedItems": "Failed Items ({count})",
|
||||||
|
"columnName": "Model Name",
|
||||||
|
"columnError": "Error",
|
||||||
|
"close": "Close",
|
||||||
|
"copyReport": "Copy Report",
|
||||||
|
"retryFailed": "Retry Failed ({count})"
|
||||||
}
|
}
|
||||||
},
|
},
|
||||||
"modelTags": {
|
"modelTags": {
|
||||||
@@ -2031,7 +2059,6 @@
|
|||||||
"presetNameTooLong": "Preset name must be {max} characters or less",
|
"presetNameTooLong": "Preset name must be {max} characters or less",
|
||||||
"presetNameInvalidChars": "Preset name contains invalid characters",
|
"presetNameInvalidChars": "Preset name contains invalid characters",
|
||||||
"presetNameExists": "A preset with this name already exists",
|
"presetNameExists": "A preset with this name already exists",
|
||||||
"maxPresetsReached": "Maximum {max} presets allowed. Delete one to add more.",
|
|
||||||
"presetNotFound": "Preset not found",
|
"presetNotFound": "Preset not found",
|
||||||
"invalidPreset": "Invalid preset data",
|
"invalidPreset": "Invalid preset data",
|
||||||
"deletePresetFailed": "Failed to delete preset",
|
"deletePresetFailed": "Failed to delete preset",
|
||||||
|
|||||||
+2249
-2222
File diff suppressed because it is too large
Load Diff
+2249
-2222
File diff suppressed because it is too large
Load Diff
+2249
-2222
File diff suppressed because it is too large
Load Diff
+2249
-2222
File diff suppressed because it is too large
Load Diff
+2249
-2222
File diff suppressed because it is too large
Load Diff
+2249
-2222
File diff suppressed because it is too large
Load Diff
+2249
-2222
File diff suppressed because it is too large
Load Diff
+2249
-2222
File diff suppressed because it is too large
Load Diff
@@ -2,7 +2,8 @@ import json
|
|||||||
import os
|
import os
|
||||||
import re
|
import re
|
||||||
|
|
||||||
from .constants import CLIP_SKIP_SENTINEL, MODELS, PROMPTS, SAMPLING, LORAS, SIZE, IMAGES, IS_SAMPLER, OVERWRITE, METADATA_OVERWRITE_FIELDS
|
from .constants import MODELS, PROMPTS, SAMPLING, LORAS, SIZE, IMAGES, IS_SAMPLER, OVERWRITE
|
||||||
|
from .overwrite_utils import collect_overwrite_params
|
||||||
|
|
||||||
|
|
||||||
def _store_checkpoint_metadata(metadata, node_id, model_name):
|
def _store_checkpoint_metadata(metadata, node_id, model_name):
|
||||||
@@ -1233,14 +1234,7 @@ class MetadataOverwriteExtractor(NodeMetadataExtractor):
|
|||||||
if not inputs:
|
if not inputs:
|
||||||
return
|
return
|
||||||
|
|
||||||
overwrite_params = {}
|
overwrite_params = collect_overwrite_params(inputs)
|
||||||
for key in METADATA_OVERWRITE_FIELDS:
|
|
||||||
value = inputs.get(key)
|
|
||||||
if key == "clip_skip":
|
|
||||||
if value != CLIP_SKIP_SENTINEL:
|
|
||||||
overwrite_params[key] = value
|
|
||||||
elif value: # truthy — only overwrite when user provided a real value
|
|
||||||
overwrite_params[key] = value
|
|
||||||
|
|
||||||
if overwrite_params:
|
if overwrite_params:
|
||||||
metadata.setdefault(OVERWRITE, {})
|
metadata.setdefault(OVERWRITE, {})
|
||||||
|
|||||||
@@ -0,0 +1,42 @@
|
|||||||
|
"""Shared helpers for Metadata Overwrite node metadata collection.
|
||||||
|
|
||||||
|
Used by both the MetadataOverwriteLM node (execution time) and the
|
||||||
|
MetadataOverwriteExtractor (hook time) so the conversion/filtering logic
|
||||||
|
cannot drift between the two paths.
|
||||||
|
"""
|
||||||
|
|
||||||
|
import logging
|
||||||
|
from typing import Any, Dict
|
||||||
|
|
||||||
|
from ..utils.utils import model_patcher_to_name
|
||||||
|
from .constants import CLIP_SKIP_SENTINEL, METADATA_OVERWRITE_FIELDS
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
|
def collect_overwrite_params(values: Dict[str, Any]) -> Dict[str, Any]:
|
||||||
|
"""Convert node input values into non-default overwrite parameters.
|
||||||
|
|
||||||
|
For most fields, a falsy value (empty string, 0) means "not set" and is
|
||||||
|
skipped. clip_skip uses a dedicated sentinel (-25) so that a wired value
|
||||||
|
of 0 is preserved. The ``model`` field accepts either a manual string or
|
||||||
|
a wired MODEL (ModelPatcher) connection; in the latter case the source
|
||||||
|
model name is extracted from the patcher's ``cached_patcher_init`` and
|
||||||
|
stored as a ComfyUI-style relative path.
|
||||||
|
"""
|
||||||
|
result: Dict[str, Any] = {}
|
||||||
|
for key in METADATA_OVERWRITE_FIELDS:
|
||||||
|
value = values.get(key)
|
||||||
|
if key == "model" and not isinstance(value, str):
|
||||||
|
value = model_patcher_to_name(value)
|
||||||
|
if value is None:
|
||||||
|
logger.warning(
|
||||||
|
"Could not extract model name from wired MODEL input "
|
||||||
|
"(no cached_patcher_init); model metadata overwrite skipped"
|
||||||
|
)
|
||||||
|
if key == "clip_skip":
|
||||||
|
if value != CLIP_SKIP_SENTINEL:
|
||||||
|
result[key] = value
|
||||||
|
elif value:
|
||||||
|
result[key] = value
|
||||||
|
return result
|
||||||
@@ -1,26 +1,102 @@
|
|||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import inspect
|
||||||
|
import re
|
||||||
|
from typing import Any
|
||||||
|
|
||||||
|
_STACK_INPUT_PATTERN = re.compile(r"^lora_stack(?:_([ab])|(\d+))$")
|
||||||
|
|
||||||
|
|
||||||
|
def _is_stack_input(name: str) -> bool:
|
||||||
|
return bool(_STACK_INPUT_PATTERN.match(name))
|
||||||
|
|
||||||
|
|
||||||
|
def _stack_slot_number(name: str) -> int:
|
||||||
|
"""Numeric slot used to order stack inputs; legacy a/b map to 1/2."""
|
||||||
|
match = _STACK_INPUT_PATTERN.match(name)
|
||||||
|
if not match:
|
||||||
|
return -1
|
||||||
|
letter, digits = match.group(1), match.group(2)
|
||||||
|
if digits is not None:
|
||||||
|
return int(digits)
|
||||||
|
return 1 if letter == "a" else 2
|
||||||
|
|
||||||
|
|
||||||
|
class _LoraStackOptionalInputs:
|
||||||
|
"""Lookup that preserves explicit optional inputs and dynamic lora_stack slots."""
|
||||||
|
|
||||||
|
def __init__(self, explicit_inputs: dict[str, tuple[str, dict[str, Any]]]) -> None:
|
||||||
|
self._explicit_inputs = explicit_inputs
|
||||||
|
|
||||||
|
def __contains__(self, item: object) -> bool:
|
||||||
|
if not isinstance(item, str):
|
||||||
|
return False
|
||||||
|
return item in self._explicit_inputs or _is_stack_input(item)
|
||||||
|
|
||||||
|
def __getitem__(self, key: str) -> tuple[str, dict[str, Any]]:
|
||||||
|
if key in self._explicit_inputs:
|
||||||
|
return self._explicit_inputs[key]
|
||||||
|
if _is_stack_input(key):
|
||||||
|
return (
|
||||||
|
"LORA_STACK",
|
||||||
|
{
|
||||||
|
"tooltip": "A LoRA stack to combine. Connect to add more inputs.",
|
||||||
|
},
|
||||||
|
)
|
||||||
|
raise KeyError(key)
|
||||||
|
|
||||||
|
|
||||||
class LoraStackCombinerLM:
|
class LoraStackCombinerLM:
|
||||||
NAME = "Lora Stack Combiner (LoraManager)"
|
NAME = "Lora Stack Combiner (LoraManager)"
|
||||||
CATEGORY = "Lora Manager/stackers"
|
CATEGORY = "Lora Manager/stackers"
|
||||||
|
DESCRIPTION = (
|
||||||
|
"Combines multiple LoRA stacks into a single stack. "
|
||||||
|
"Supports dynamic inputs: connect a stack to add more inputs."
|
||||||
|
)
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def INPUT_TYPES(cls):
|
def INPUT_TYPES(cls):
|
||||||
|
optional_inputs: dict[str, tuple[str, dict[str, Any]]] = {
|
||||||
|
"lora_stack1": (
|
||||||
|
"LORA_STACK",
|
||||||
|
{
|
||||||
|
"tooltip": "A LoRA stack to combine. Connect to add more inputs.",
|
||||||
|
},
|
||||||
|
),
|
||||||
|
"lora_stack2": (
|
||||||
|
"LORA_STACK",
|
||||||
|
{
|
||||||
|
"tooltip": "A LoRA stack to combine. Connect to add more inputs.",
|
||||||
|
},
|
||||||
|
),
|
||||||
|
}
|
||||||
|
|
||||||
|
stack = inspect.stack()
|
||||||
|
if len(stack) > 2 and stack[2].function == "get_input_info":
|
||||||
|
optional_inputs = _LoraStackOptionalInputs(optional_inputs) # type: ignore[assignment]
|
||||||
|
|
||||||
return {
|
return {
|
||||||
"required": {
|
"required": {},
|
||||||
"lora_stack_a": ("LORA_STACK",),
|
"optional": optional_inputs,
|
||||||
"lora_stack_b": ("LORA_STACK",),
|
|
||||||
},
|
|
||||||
}
|
}
|
||||||
|
|
||||||
RETURN_TYPES = ("LORA_STACK",)
|
RETURN_TYPES = ("LORA_STACK",)
|
||||||
RETURN_NAMES = ("LORA_STACK",)
|
RETURN_NAMES = ("LORA_STACK",)
|
||||||
FUNCTION = "combine_stacks"
|
FUNCTION = "combine_stacks"
|
||||||
|
|
||||||
def combine_stacks(self, lora_stack_a, lora_stack_b):
|
def combine_stacks(self, lora_stack1=None, lora_stack2=None, **kwargs):
|
||||||
combined_stack = []
|
stacks = {
|
||||||
|
"lora_stack1": lora_stack1,
|
||||||
|
"lora_stack2": lora_stack2,
|
||||||
|
}
|
||||||
|
for key, value in kwargs.items():
|
||||||
|
if _is_stack_input(key) and value is not None:
|
||||||
|
stacks[key] = value
|
||||||
|
|
||||||
if lora_stack_a:
|
combined_stack = []
|
||||||
combined_stack.extend(lora_stack_a)
|
for key in sorted(stacks, key=_stack_slot_number):
|
||||||
if lora_stack_b:
|
stack = stacks[key]
|
||||||
combined_stack.extend(lora_stack_b)
|
if stack:
|
||||||
|
combined_stack.extend(stack)
|
||||||
|
|
||||||
return (combined_stack,)
|
return (combined_stack,)
|
||||||
|
|||||||
@@ -9,10 +9,8 @@ but users may wire 0 to express "no clip skip / default".
|
|||||||
|
|
||||||
from typing import Any
|
from typing import Any
|
||||||
|
|
||||||
from ..metadata_collector.constants import (
|
from ..metadata_collector.constants import CLIP_SKIP_SENTINEL as _CLIP_SKIP_SENTINEL
|
||||||
CLIP_SKIP_SENTINEL as _CLIP_SKIP_SENTINEL,
|
from ..metadata_collector.overwrite_utils import collect_overwrite_params
|
||||||
METADATA_OVERWRITE_FIELDS,
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
class MetadataOverwriteLM:
|
class MetadataOverwriteLM:
|
||||||
@@ -87,12 +85,16 @@ class MetadataOverwriteLM:
|
|||||||
},
|
},
|
||||||
),
|
),
|
||||||
"model": (
|
"model": (
|
||||||
"STRING",
|
"STRING,MODEL",
|
||||||
{
|
{
|
||||||
"default": "",
|
"default": "",
|
||||||
|
"widgetType": "STRING",
|
||||||
"tooltip": (
|
"tooltip": (
|
||||||
"The checkpoint or diffusion model (UNet) used "
|
"The checkpoint or diffusion model (UNet) used "
|
||||||
"for generation. Only overwrites when non-empty."
|
"for generation. Fill in the name manually or "
|
||||||
|
"connect a MODEL output — the model name is then "
|
||||||
|
"extracted automatically. Only overwrites when "
|
||||||
|
"non-empty."
|
||||||
),
|
),
|
||||||
},
|
},
|
||||||
),
|
),
|
||||||
@@ -158,13 +160,10 @@ class MetadataOverwriteLM:
|
|||||||
For most fields, a falsy value (empty string, 0) means "not set"
|
For most fields, a falsy value (empty string, 0) means "not set"
|
||||||
and is skipped. clip_skip uses a dedicated sentinel (-25) so that
|
and is skipped. clip_skip uses a dedicated sentinel (-25) so that
|
||||||
a wired value of 0 is preserved and reaches the metadata pipeline.
|
a wired value of 0 is preserved and reaches the metadata pipeline.
|
||||||
|
|
||||||
|
The ``model`` field accepts either a manual string or a wired MODEL
|
||||||
|
(ModelPatcher) connection; in the latter case the underlying model
|
||||||
|
name is extracted from the patcher's ``cached_patcher_init`` and
|
||||||
|
stored as a ComfyUI-style relative path.
|
||||||
"""
|
"""
|
||||||
result: dict[str, Any] = {}
|
return (collect_overwrite_params(kwargs),)
|
||||||
for key in METADATA_OVERWRITE_FIELDS:
|
|
||||||
value = kwargs.get(key)
|
|
||||||
if key == "clip_skip":
|
|
||||||
if value != _CLIP_SKIP_SENTINEL:
|
|
||||||
result[key] = value
|
|
||||||
elif value:
|
|
||||||
result[key] = value
|
|
||||||
return (result,)
|
|
||||||
|
|||||||
+16
-3
@@ -252,6 +252,13 @@ class SaveImageLM:
|
|||||||
"tooltip": "When enabled, embeds generation parameters into the saved image metadata. Disable to skip writing generation metadata.",
|
"tooltip": "When enabled, embeds generation parameters into the saved image metadata. Disable to skip writing generation metadata.",
|
||||||
},
|
},
|
||||||
),
|
),
|
||||||
|
"add_loras_to_prompt": (
|
||||||
|
"BOOLEAN",
|
||||||
|
{
|
||||||
|
"default": False,
|
||||||
|
"tooltip": "When enabled, appends the LoRA syntax line (e.g. <lora:name:strength>) after the positive prompt in the saved metadata.",
|
||||||
|
},
|
||||||
|
),
|
||||||
"add_counter_to_filename": (
|
"add_counter_to_filename": (
|
||||||
"BOOLEAN",
|
"BOOLEAN",
|
||||||
{
|
{
|
||||||
@@ -348,7 +355,7 @@ class SaveImageLM:
|
|||||||
type_lower = model_type.lower() if model_type else "other"
|
type_lower = model_type.lower() if model_type else "other"
|
||||||
return f"urn:air:{slug}:{type_lower}:civitai:{model_id}@{version_id}"
|
return f"urn:air:{slug}:{type_lower}:civitai:{model_id}@{version_id}"
|
||||||
|
|
||||||
def format_metadata(self, metadata_dict: dict) -> str:
|
def format_metadata(self, metadata_dict: dict, add_loras_to_prompt: bool = False) -> str:
|
||||||
"""Format metadata as A1111-compatible parameters string with Hashes JSON and Civitai resources."""
|
"""Format metadata as A1111-compatible parameters string with Hashes JSON and Civitai resources."""
|
||||||
if not metadata_dict: return ""
|
if not metadata_dict: return ""
|
||||||
|
|
||||||
@@ -458,7 +465,10 @@ class SaveImageLM:
|
|||||||
scheduler_name = scheduler_mapping.get(scheduler, scheduler) if scheduler else None
|
scheduler_name = scheduler_mapping.get(scheduler, scheduler) if scheduler else None
|
||||||
|
|
||||||
# Build output lines
|
# Build output lines
|
||||||
lines = [prompt] if prompt else [""]
|
prompt_line = prompt if prompt else ""
|
||||||
|
if add_loras_to_prompt and loras_text:
|
||||||
|
prompt_line = f"{prompt_line}\n{loras_text}" if prompt_line else loras_text
|
||||||
|
lines = [prompt_line] if prompt_line else [""]
|
||||||
if negative_prompt:
|
if negative_prompt:
|
||||||
lines.append(f"Negative prompt: {negative_prompt}")
|
lines.append(f"Negative prompt: {negative_prompt}")
|
||||||
|
|
||||||
@@ -793,6 +803,7 @@ class SaveImageLM:
|
|||||||
save_with_metadata=True,
|
save_with_metadata=True,
|
||||||
add_counter_to_filename=True,
|
add_counter_to_filename=True,
|
||||||
save_as_recipe=False,
|
save_as_recipe=False,
|
||||||
|
add_loras_to_prompt=False,
|
||||||
):
|
):
|
||||||
"""Save images with metadata"""
|
"""Save images with metadata"""
|
||||||
results = []
|
results = []
|
||||||
@@ -801,7 +812,7 @@ class SaveImageLM:
|
|||||||
raw_metadata = get_metadata()
|
raw_metadata = get_metadata()
|
||||||
metadata_dict = MetadataProcessor.to_dict(raw_metadata, id)
|
metadata_dict = MetadataProcessor.to_dict(raw_metadata, id)
|
||||||
|
|
||||||
metadata = self.format_metadata(metadata_dict)
|
metadata = self.format_metadata(metadata_dict, add_loras_to_prompt)
|
||||||
|
|
||||||
# Process filename_prefix with pattern substitution
|
# Process filename_prefix with pattern substitution
|
||||||
filename_prefix = self.format_filename(filename_prefix, metadata_dict)
|
filename_prefix = self.format_filename(filename_prefix, metadata_dict)
|
||||||
@@ -943,6 +954,7 @@ class SaveImageLM:
|
|||||||
save_with_metadata=True,
|
save_with_metadata=True,
|
||||||
add_counter_to_filename=True,
|
add_counter_to_filename=True,
|
||||||
save_as_recipe=False,
|
save_as_recipe=False,
|
||||||
|
add_loras_to_prompt=False,
|
||||||
):
|
):
|
||||||
"""Process and save image with metadata"""
|
"""Process and save image with metadata"""
|
||||||
# Make sure the output directory exists
|
# Make sure the output directory exists
|
||||||
@@ -974,6 +986,7 @@ class SaveImageLM:
|
|||||||
save_with_metadata,
|
save_with_metadata,
|
||||||
add_counter_to_filename,
|
add_counter_to_filename,
|
||||||
save_as_recipe,
|
save_as_recipe,
|
||||||
|
add_loras_to_prompt,
|
||||||
)
|
)
|
||||||
|
|
||||||
return {
|
return {
|
||||||
|
|||||||
@@ -7,6 +7,21 @@ from ..utils.utils import get_checkpoint_info_absolute, _format_model_name_for_c
|
|||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
|
def _reload_gguf_unet(
|
||||||
|
unet_path: str, weight_dtype: str, disable_dynamic: bool = False
|
||||||
|
) -> object:
|
||||||
|
"""Reload a GGUF diffusion model from disk (cached_patcher_init factory).
|
||||||
|
|
||||||
|
Mirrors the GGUF branch of UNETLoaderLM.load_unet so ModelPatcher
|
||||||
|
deepclone/dynamic machinery can rebuild GGUF models with the correct
|
||||||
|
GGMLOps. ``disable_dynamic`` is accepted for signature compatibility
|
||||||
|
with core ComfyUI loaders.
|
||||||
|
"""
|
||||||
|
loader = UNETLoaderLM()
|
||||||
|
model, = loader._load_gguf_unet(unet_path, unet_path, weight_dtype)
|
||||||
|
return model
|
||||||
|
|
||||||
|
|
||||||
class UNETLoaderLM:
|
class UNETLoaderLM:
|
||||||
"""UNET Loader with support for extra folder paths
|
"""UNET Loader with support for extra folder paths
|
||||||
|
|
||||||
@@ -196,6 +211,12 @@ class UNETLoaderLM:
|
|||||||
# Wrap with GGUFModelPatcher
|
# Wrap with GGUFModelPatcher
|
||||||
model = GGUFModelPatcher.clone(model)
|
model = GGUFModelPatcher.clone(model)
|
||||||
|
|
||||||
|
# Register a reload factory so the MODEL carries its source path
|
||||||
|
# (cached_patcher_init) like core ComfyUI loaders do — required
|
||||||
|
# for model-name extraction downstream and for ModelPatcher
|
||||||
|
# deepclone/dynamic machinery.
|
||||||
|
model.cached_patcher_init = (_reload_gguf_unet, (unet_path, weight_dtype))
|
||||||
|
|
||||||
return (model,)
|
return (model,)
|
||||||
|
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
|
|||||||
@@ -2590,6 +2590,8 @@ class ModelLibraryHandler:
|
|||||||
status=400,
|
status=400,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
cursor = request.query.get("cursor")
|
||||||
|
|
||||||
metadata_provider = await self._metadata_provider_factory()
|
metadata_provider = await self._metadata_provider_factory()
|
||||||
if not metadata_provider:
|
if not metadata_provider:
|
||||||
return web.json_response(
|
return web.json_response(
|
||||||
@@ -2598,7 +2600,7 @@ class ModelLibraryHandler:
|
|||||||
)
|
)
|
||||||
|
|
||||||
try:
|
try:
|
||||||
models = await metadata_provider.get_user_models(username)
|
result = await metadata_provider.get_user_models(username, cursor)
|
||||||
except NotImplementedError:
|
except NotImplementedError:
|
||||||
return web.json_response(
|
return web.json_response(
|
||||||
{
|
{
|
||||||
@@ -2608,14 +2610,35 @@ class ModelLibraryHandler:
|
|||||||
status=501,
|
status=501,
|
||||||
)
|
)
|
||||||
|
|
||||||
if models is None:
|
if result is None:
|
||||||
return web.json_response(
|
return web.json_response(
|
||||||
{"success": False, "error": "Failed to fetch user models"},
|
{"success": False, "error": "Failed to fetch user models"},
|
||||||
status=502,
|
status=502,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
if isinstance(result, dict):
|
||||||
|
models = result.get("items")
|
||||||
|
next_cursor = result.get("nextCursor")
|
||||||
|
else:
|
||||||
|
# Defensive: tolerate providers that still return a raw list
|
||||||
|
models = result
|
||||||
|
next_cursor = None
|
||||||
|
|
||||||
if not isinstance(models, list):
|
if not isinstance(models, list):
|
||||||
models = []
|
models = []
|
||||||
|
if next_cursor is not None and not isinstance(next_cursor, str):
|
||||||
|
next_cursor = str(next_cursor)
|
||||||
|
|
||||||
|
estimated_total = None
|
||||||
|
if cursor is None:
|
||||||
|
get_count = getattr(metadata_provider, "get_creator_model_count", None)
|
||||||
|
if get_count is not None:
|
||||||
|
try:
|
||||||
|
estimated_total = await get_count(username)
|
||||||
|
except Exception: # best-effort only
|
||||||
|
estimated_total = None
|
||||||
|
if not isinstance(estimated_total, int):
|
||||||
|
estimated_total = None
|
||||||
|
|
||||||
lora_scanner = await self._service_registry.get_lora_scanner()
|
lora_scanner = await self._service_registry.get_lora_scanner()
|
||||||
checkpoint_scanner = await self._service_registry.get_checkpoint_scanner()
|
checkpoint_scanner = await self._service_registry.get_checkpoint_scanner()
|
||||||
@@ -2635,6 +2658,7 @@ class ModelLibraryHandler:
|
|||||||
versions: list[dict] = []
|
versions: list[dict] = []
|
||||||
history_service = await self._get_download_history_service()
|
history_service = await self._get_download_history_service()
|
||||||
model_ids: list[int] = []
|
model_ids: list[int] = []
|
||||||
|
model_count = 0
|
||||||
for model in models:
|
for model in models:
|
||||||
try:
|
try:
|
||||||
model_ids.append(int(model.get("id")))
|
model_ids.append(int(model.get("id")))
|
||||||
@@ -2668,6 +2692,8 @@ class ModelLibraryHandler:
|
|||||||
if model_type not in normalized_allowed_types:
|
if model_type not in normalized_allowed_types:
|
||||||
continue
|
continue
|
||||||
|
|
||||||
|
model_count += 1
|
||||||
|
|
||||||
scanner = type_scanner_map.get(model_type)
|
scanner = type_scanner_map.get(model_type)
|
||||||
if scanner is None:
|
if scanner is None:
|
||||||
return web.json_response(
|
return web.json_response(
|
||||||
@@ -2733,7 +2759,15 @@ class ModelLibraryHandler:
|
|||||||
)
|
)
|
||||||
|
|
||||||
return web.json_response(
|
return web.json_response(
|
||||||
{"success": True, "username": username, "versions": versions}
|
{
|
||||||
|
"success": True,
|
||||||
|
"username": username,
|
||||||
|
"versions": versions,
|
||||||
|
"modelCount": model_count,
|
||||||
|
"nextCursor": next_cursor,
|
||||||
|
"hasMore": next_cursor is not None,
|
||||||
|
"estimatedTotal": estimated_total,
|
||||||
|
}
|
||||||
)
|
)
|
||||||
except Exception as exc: # pragma: no cover - defensive logging
|
except Exception as exc: # pragma: no cover - defensive logging
|
||||||
logger.error("Failed to get Civitai user models: %s", exc, exc_info=True)
|
logger.error("Failed to get Civitai user models: %s", exc, exc_info=True)
|
||||||
|
|||||||
@@ -1488,8 +1488,73 @@ class ModelQueryHandler:
|
|||||||
search = request.query.get("search", "").strip()
|
search = request.query.get("search", "").strip()
|
||||||
limit = min(int(request.query.get("limit", "15")), 100)
|
limit = min(int(request.query.get("limit", "15")), 100)
|
||||||
offset = max(0, int(request.query.get("offset", "0")))
|
offset = max(0, int(request.query.get("offset", "0")))
|
||||||
|
|
||||||
|
folder = request.query.get("folder")
|
||||||
|
recursive = request.query.get("recursive", "true").lower() == "true"
|
||||||
|
base_models = list(request.query.getall("base_model", []))
|
||||||
|
model_types = list(request.query.getall("model_type", []))
|
||||||
|
|
||||||
|
tag_filters: Dict[str, str] = {}
|
||||||
|
for tag in request.query.getall("tag_include", []):
|
||||||
|
if tag:
|
||||||
|
tag_filters[tag] = "include"
|
||||||
|
for tag in request.query.getall("tag_exclude", []):
|
||||||
|
if tag:
|
||||||
|
tag_filters[tag] = "exclude"
|
||||||
|
|
||||||
|
auto_tag_filters: Dict[str, str] = {}
|
||||||
|
for tag in request.query.getall("auto_tag_include", []):
|
||||||
|
if tag:
|
||||||
|
auto_tag_filters[tag] = "include"
|
||||||
|
for tag in request.query.getall("auto_tag_exclude", []):
|
||||||
|
if tag:
|
||||||
|
auto_tag_filters[tag] = "exclude"
|
||||||
|
|
||||||
|
tag_logic = request.query.get("tag_logic", "any").lower()
|
||||||
|
if tag_logic not in ("any", "all"):
|
||||||
|
tag_logic = "any"
|
||||||
|
|
||||||
|
credit_required = request.query.get("credit_required")
|
||||||
|
if credit_required is not None:
|
||||||
|
credit_required = credit_required.lower() not in ("false", "0", "")
|
||||||
|
|
||||||
|
allow_selling_generated_content = request.query.get(
|
||||||
|
"allow_selling_generated_content"
|
||||||
|
)
|
||||||
|
if allow_selling_generated_content is not None:
|
||||||
|
allow_selling_generated_content = (
|
||||||
|
allow_selling_generated_content.lower() not in ("false", "0", "")
|
||||||
|
)
|
||||||
|
|
||||||
|
# The presence of the recursive param (always sent by the loras
|
||||||
|
# widget when filter mode is on) signals that the filter pipeline
|
||||||
|
# must run even when no concrete filter is set, so global settings
|
||||||
|
# like show_only_sfw stay consistent with the list endpoint.
|
||||||
|
apply_filters = (
|
||||||
|
"recursive" in request.query
|
||||||
|
or folder is not None
|
||||||
|
or bool(base_models)
|
||||||
|
or bool(model_types)
|
||||||
|
or bool(tag_filters)
|
||||||
|
or bool(auto_tag_filters)
|
||||||
|
or credit_required is not None
|
||||||
|
or allow_selling_generated_content is not None
|
||||||
|
)
|
||||||
|
|
||||||
matching_paths = await self._service.search_relative_paths(
|
matching_paths = await self._service.search_relative_paths(
|
||||||
search, limit, offset
|
search,
|
||||||
|
limit,
|
||||||
|
offset,
|
||||||
|
folder=folder,
|
||||||
|
recursive=recursive,
|
||||||
|
base_models=base_models,
|
||||||
|
model_types=model_types,
|
||||||
|
tags=tag_filters,
|
||||||
|
auto_tags=auto_tag_filters,
|
||||||
|
tag_logic=tag_logic,
|
||||||
|
credit_required=credit_required,
|
||||||
|
allow_selling_generated_content=allow_selling_generated_content,
|
||||||
|
apply_filters=apply_filters,
|
||||||
)
|
)
|
||||||
return web.json_response(
|
return web.json_response(
|
||||||
{"success": True, "relative_paths": matching_paths}
|
{"success": True, "relative_paths": matching_paths}
|
||||||
|
|||||||
@@ -10,7 +10,7 @@ import asyncio
|
|||||||
import tempfile
|
import tempfile
|
||||||
from dataclasses import dataclass
|
from dataclasses import dataclass
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
from typing import Any, Awaitable, Callable, Dict, List, Mapping, Optional
|
from typing import Any, Awaitable, Callable, Dict, List, Mapping, Optional, Tuple
|
||||||
|
|
||||||
from aiohttp import web
|
from aiohttp import web
|
||||||
|
|
||||||
@@ -44,6 +44,22 @@ EnsureDependenciesCallable = Callable[[], Awaitable[None]]
|
|||||||
RecipeScannerGetter = Callable[[], Any]
|
RecipeScannerGetter = Callable[[], Any]
|
||||||
CivitaiClientGetter = Callable[[], Any]
|
CivitaiClientGetter = Callable[[], Any]
|
||||||
|
|
||||||
|
# Cap concurrent preview-dimension reads across requests. With a cold LRU
|
||||||
|
# cache one page can touch up to page_size image files; 16 balances SSD and
|
||||||
|
# HDD throughput without starving the event loop.
|
||||||
|
_DIMS_READ_SEMAPHORE = asyncio.Semaphore(16)
|
||||||
|
|
||||||
|
|
||||||
|
async def _read_preview_dims(path: str) -> Optional[Tuple[int, int]]:
|
||||||
|
"""Read preview dimensions off the event loop under the concurrency cap.
|
||||||
|
|
||||||
|
PIL I/O runs in a worker thread so it never blocks the event loop, and the
|
||||||
|
semaphore bounds how many files are opened at once even when many list
|
||||||
|
requests land together.
|
||||||
|
"""
|
||||||
|
async with _DIMS_READ_SEMAPHORE:
|
||||||
|
return await asyncio.to_thread(ExifUtils.get_image_dimensions, path)
|
||||||
|
|
||||||
|
|
||||||
@dataclass(frozen=True)
|
@dataclass(frozen=True)
|
||||||
class RecipeHandlerSet:
|
class RecipeHandlerSet:
|
||||||
@@ -246,7 +262,8 @@ class RecipeListingHandler:
|
|||||||
recursive=recursive,
|
recursive=recursive,
|
||||||
)
|
)
|
||||||
|
|
||||||
for item in result.get("items", []):
|
items = result.get("items", [])
|
||||||
|
for item in items:
|
||||||
file_path = item.get("file_path")
|
file_path = item.get("file_path")
|
||||||
if file_path:
|
if file_path:
|
||||||
item["file_url"] = self.format_recipe_file_url(file_path)
|
item["file_url"] = self.format_recipe_file_url(file_path)
|
||||||
@@ -255,6 +272,26 @@ class RecipeListingHandler:
|
|||||||
item.setdefault("loras", [])
|
item.setdefault("loras", [])
|
||||||
item.setdefault("base_model", "")
|
item.setdefault("base_model", "")
|
||||||
|
|
||||||
|
# Batch preview dimension reads with asyncio.gather. The previous
|
||||||
|
# loop awaited asyncio.to_thread once per item, so a page_size=100
|
||||||
|
# request submitted 100 sequential thread calls (50-300ms cold-page
|
||||||
|
# latency). gather runs them concurrently while the semaphore caps
|
||||||
|
# disk opens; dimensions stay omitted (not null) when a preview has
|
||||||
|
# no readable size (video, missing file).
|
||||||
|
to_read = [
|
||||||
|
(i, item.get("file_path"))
|
||||||
|
for i, item in enumerate(items)
|
||||||
|
if item.get("file_path")
|
||||||
|
]
|
||||||
|
if to_read:
|
||||||
|
dims_list = await asyncio.gather(
|
||||||
|
*(_read_preview_dims(path) for _, path in to_read)
|
||||||
|
)
|
||||||
|
for (idx, _), dims in zip(to_read, dims_list):
|
||||||
|
if dims:
|
||||||
|
item = items[idx]
|
||||||
|
item["width"], item["height"] = dims
|
||||||
|
|
||||||
return web.json_response(result)
|
return web.json_response(result)
|
||||||
except Exception as exc:
|
except Exception as exc:
|
||||||
self._logger.error("Error retrieving recipes: %s", exc, exc_info=True)
|
self._logger.error("Error retrieving recipes: %s", exc, exc_info=True)
|
||||||
|
|||||||
@@ -1,6 +1,7 @@
|
|||||||
from abc import ABC, abstractmethod
|
from abc import ABC, abstractmethod
|
||||||
import asyncio
|
import asyncio
|
||||||
import re
|
import re
|
||||||
|
import random
|
||||||
from typing import Any, Dict, List, Optional, Type, Union, TYPE_CHECKING
|
from typing import Any, Dict, List, Optional, Type, Union, TYPE_CHECKING
|
||||||
import logging
|
import logging
|
||||||
import os
|
import os
|
||||||
@@ -390,6 +391,12 @@ class BaseModelService(ABC):
|
|||||||
(item.get("model_name") or item.get("file_name") or "").lower(),
|
(item.get("model_name") or item.get("file_name") or "").lower(),
|
||||||
item.get("file_path", "").lower(),
|
item.get("file_path", "").lower(),
|
||||||
)
|
)
|
||||||
|
elif key_name == "random":
|
||||||
|
# Seeded random shuffle: same seed -> same order (stable pagination)
|
||||||
|
rng = random.Random(sort_params.seed or "random")
|
||||||
|
result = list(data)
|
||||||
|
rng.shuffle(result)
|
||||||
|
return result
|
||||||
elif key_name == "size":
|
elif key_name == "size":
|
||||||
key_fn = lambda item: (
|
key_fn = lambda item: (
|
||||||
int(item.get("size", 0) or 0),
|
int(item.get("size", 0) or 0),
|
||||||
@@ -1252,19 +1259,87 @@ class BaseModelService(ABC):
|
|||||||
)
|
)
|
||||||
|
|
||||||
async def search_relative_paths(
|
async def search_relative_paths(
|
||||||
self, search_term: str, limit: int = 15, offset: int = 0
|
self,
|
||||||
|
search_term: str,
|
||||||
|
limit: int = 15,
|
||||||
|
offset: int = 0,
|
||||||
|
*,
|
||||||
|
folder: Optional[str] = None,
|
||||||
|
folder_include: Optional[list] = None,
|
||||||
|
folder_exclude: Optional[list] = None,
|
||||||
|
base_models: Optional[list] = None,
|
||||||
|
model_types: Optional[list] = None,
|
||||||
|
tags: Optional[dict] = None,
|
||||||
|
auto_tags: Optional[dict] = None,
|
||||||
|
tag_logic: str = "any",
|
||||||
|
credit_required: Optional[bool] = None,
|
||||||
|
allow_selling_generated_content: Optional[bool] = None,
|
||||||
|
recursive: bool = True,
|
||||||
|
apply_filters: bool = False,
|
||||||
) -> List[str]:
|
) -> List[str]:
|
||||||
"""Search model relative file paths for autocomplete functionality"""
|
"""Search model relative file paths for autocomplete functionality.
|
||||||
|
|
||||||
|
Optional filter kwargs mirror the filters used by the list endpoint
|
||||||
|
(/api/lm/{prefix}/list). When no filter kwargs are provided the
|
||||||
|
behavior is identical to plain token-based path matching.
|
||||||
|
"""
|
||||||
cache = await self.scanner.get_cached_data()
|
cache = await self.scanner.get_cached_data()
|
||||||
include_terms, exclude_terms = self._parse_search_tokens(search_term)
|
include_terms, exclude_terms = self._parse_search_tokens(search_term)
|
||||||
|
|
||||||
|
data = cache.raw_data
|
||||||
|
has_filters = any(
|
||||||
|
[
|
||||||
|
apply_filters,
|
||||||
|
folder is not None,
|
||||||
|
folder_include,
|
||||||
|
folder_exclude,
|
||||||
|
base_models,
|
||||||
|
model_types,
|
||||||
|
tags,
|
||||||
|
auto_tags,
|
||||||
|
credit_required is not None,
|
||||||
|
allow_selling_generated_content is not None,
|
||||||
|
]
|
||||||
|
)
|
||||||
|
if has_filters:
|
||||||
|
# Auto-tags are not stored in the scanner cache — they are computed
|
||||||
|
# on the fly. Pre-compute them only when an auto-tag filter is
|
||||||
|
# active to avoid mutating cache entries unnecessarily.
|
||||||
|
if auto_tags:
|
||||||
|
from .auto_tag_service import extract_auto_tags
|
||||||
|
|
||||||
|
for item in data:
|
||||||
|
if not item.get("auto_tags"):
|
||||||
|
item["auto_tags"] = extract_auto_tags(item)
|
||||||
|
|
||||||
|
criteria = FilterCriteria(
|
||||||
|
folder=folder,
|
||||||
|
folder_include=folder_include,
|
||||||
|
folder_exclude=folder_exclude,
|
||||||
|
base_models=base_models,
|
||||||
|
model_types=model_types,
|
||||||
|
tags=tags,
|
||||||
|
auto_tags=auto_tags,
|
||||||
|
search_options={"recursive": recursive},
|
||||||
|
tag_logic=tag_logic,
|
||||||
|
)
|
||||||
|
data = self.filter_set.apply(data, criteria)
|
||||||
|
if credit_required is not None:
|
||||||
|
data = await self._apply_credit_required_filter(
|
||||||
|
data, credit_required
|
||||||
|
)
|
||||||
|
if allow_selling_generated_content is not None:
|
||||||
|
data = await self._apply_allow_selling_filter(
|
||||||
|
data, allow_selling_generated_content
|
||||||
|
)
|
||||||
|
|
||||||
matching_paths = []
|
matching_paths = []
|
||||||
|
|
||||||
# Get model roots for path calculation
|
# Get model roots for path calculation
|
||||||
model_roots = self.scanner.get_model_roots()
|
model_roots = self.scanner.get_model_roots()
|
||||||
|
|
||||||
# Collect all matching paths first (needed for proper sorting and offset)
|
# Collect all matching paths first (needed for proper sorting and offset)
|
||||||
for model in cache.raw_data:
|
for model in data:
|
||||||
file_path = model.get("file_path", "")
|
file_path = model.get("file_path", "")
|
||||||
if not file_path:
|
if not file_path:
|
||||||
continue
|
continue
|
||||||
|
|||||||
@@ -2,6 +2,7 @@ import asyncio
|
|||||||
import copy
|
import copy
|
||||||
import logging
|
import logging
|
||||||
import os
|
import os
|
||||||
|
import time
|
||||||
from collections import OrderedDict
|
from collections import OrderedDict
|
||||||
from typing import Any, Optional, Dict, Tuple, List, Sequence
|
from typing import Any, Optional, Dict, Tuple, List, Sequence
|
||||||
from .connectivity_guard import (
|
from .connectivity_guard import (
|
||||||
@@ -19,6 +20,12 @@ from ..utils.civitai_utils import resolve_license_payload
|
|||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
# Best-effort cache for creator model counts, keyed by lowercase username.
|
||||||
|
# Values are (monotonic timestamp, count or None); None results are cached
|
||||||
|
# too so repeated failures don't hammer the API.
|
||||||
|
_CREATOR_COUNT_CACHE_TTL_SECONDS = 600
|
||||||
|
_creator_model_count_cache: Dict[str, Tuple[float, Optional[int]]] = {}
|
||||||
|
|
||||||
|
|
||||||
class CivitaiClient:
|
class CivitaiClient:
|
||||||
_instance = None
|
_instance = None
|
||||||
@@ -743,17 +750,34 @@ class CivitaiClient:
|
|||||||
|
|
||||||
return all_versions if all_versions else None
|
return all_versions if all_versions else None
|
||||||
|
|
||||||
async def get_user_models(self, username: str) -> Optional[List[Dict]]:
|
async def get_user_models(
|
||||||
"""Fetch all models for a specific Civitai user."""
|
self, username: str, cursor: Optional[str] = None
|
||||||
|
) -> Optional[Dict[str, Any]]:
|
||||||
|
"""Fetch one page (up to 100 models) for a specific Civitai user.
|
||||||
|
|
||||||
|
Returns ``{"items": [...], "nextCursor": <str|None>}`` on success,
|
||||||
|
or None on failure. Pass ``cursor`` (from a previous response's
|
||||||
|
``nextCursor``) to fetch subsequent pages.
|
||||||
|
"""
|
||||||
if not username:
|
if not username:
|
||||||
return None
|
return None
|
||||||
|
|
||||||
|
params: Dict[str, Any] = {
|
||||||
|
"username": username,
|
||||||
|
"nsfw": "true",
|
||||||
|
"limit": 100,
|
||||||
|
"sort": "Newest",
|
||||||
|
"period": "AllTime",
|
||||||
|
}
|
||||||
|
if cursor:
|
||||||
|
params["cursor"] = cursor
|
||||||
|
|
||||||
try:
|
try:
|
||||||
success, result = await self._make_request(
|
success, result = await self._make_request(
|
||||||
"GET",
|
"GET",
|
||||||
f"{self.base_url}/models",
|
f"{self.base_url}/models",
|
||||||
use_auth=True,
|
use_auth=True,
|
||||||
params={"username": username, "nsfw": "true"},
|
params=params,
|
||||||
)
|
)
|
||||||
|
|
||||||
if not success:
|
if not success:
|
||||||
@@ -765,7 +789,7 @@ class CivitaiClient:
|
|||||||
|
|
||||||
items = result.get("items") if isinstance(result, dict) else None
|
items = result.get("items") if isinstance(result, dict) else None
|
||||||
if not isinstance(items, list):
|
if not isinstance(items, list):
|
||||||
return []
|
items = []
|
||||||
|
|
||||||
for model in items:
|
for model in items:
|
||||||
versions = model.get("modelVersions")
|
versions = model.get("modelVersions")
|
||||||
@@ -774,9 +798,68 @@ class CivitaiClient:
|
|||||||
for version in versions:
|
for version in versions:
|
||||||
self._remove_comfy_metadata(version)
|
self._remove_comfy_metadata(version)
|
||||||
|
|
||||||
return items
|
next_cursor: Optional[str] = None
|
||||||
|
metadata = result.get("metadata") if isinstance(result, dict) else None
|
||||||
|
if isinstance(metadata, dict):
|
||||||
|
raw_cursor = metadata.get("nextCursor")
|
||||||
|
if raw_cursor is not None:
|
||||||
|
next_cursor = str(raw_cursor)
|
||||||
|
|
||||||
|
return {"items": items, "nextCursor": next_cursor}
|
||||||
except RateLimitError:
|
except RateLimitError:
|
||||||
raise
|
raise
|
||||||
except Exception as exc: # pragma: no cover - defensive logging
|
except Exception as exc: # pragma: no cover - defensive logging
|
||||||
logger.error("Error fetching models for %s: %s", username, exc)
|
logger.error("Error fetching models for %s: %s", username, exc)
|
||||||
return None
|
return None
|
||||||
|
|
||||||
|
async def get_creator_model_count(self, username: str) -> Optional[int]:
|
||||||
|
"""Best-effort lookup of a creator's published model count.
|
||||||
|
|
||||||
|
Uses the ``/creators`` endpoint (a contains-match query), picking the
|
||||||
|
entry whose username matches exactly (case-insensitive). Returns None
|
||||||
|
on any failure; never raises. Results (including None) are cached
|
||||||
|
for ``_CREATOR_COUNT_CACHE_TTL_SECONDS``.
|
||||||
|
"""
|
||||||
|
if not username:
|
||||||
|
return None
|
||||||
|
|
||||||
|
cache_key = username.lower()
|
||||||
|
cached = _creator_model_count_cache.get(cache_key)
|
||||||
|
if cached is not None:
|
||||||
|
cached_at, cached_count = cached
|
||||||
|
if time.monotonic() - cached_at < _CREATOR_COUNT_CACHE_TTL_SECONDS:
|
||||||
|
return cached_count
|
||||||
|
|
||||||
|
count: Optional[int] = None
|
||||||
|
try:
|
||||||
|
success, result = await self._make_request(
|
||||||
|
"GET",
|
||||||
|
f"{self.base_url}/creators",
|
||||||
|
use_auth=True,
|
||||||
|
params={"query": username, "limit": 10},
|
||||||
|
)
|
||||||
|
|
||||||
|
if success and isinstance(result, dict):
|
||||||
|
creators = result.get("items")
|
||||||
|
if isinstance(creators, list):
|
||||||
|
for creator in creators:
|
||||||
|
if not isinstance(creator, dict):
|
||||||
|
continue
|
||||||
|
creator_name = creator.get("username")
|
||||||
|
if not isinstance(creator_name, str):
|
||||||
|
continue
|
||||||
|
if creator_name.lower() != cache_key:
|
||||||
|
continue
|
||||||
|
model_count = creator.get("modelCount")
|
||||||
|
if isinstance(model_count, (int, float)) and not isinstance(
|
||||||
|
model_count, bool
|
||||||
|
):
|
||||||
|
count = int(model_count)
|
||||||
|
break
|
||||||
|
except Exception as exc: # best-effort only, never propagate
|
||||||
|
logger.debug(
|
||||||
|
"Failed to fetch creator model count for %s: %s", username, exc
|
||||||
|
)
|
||||||
|
|
||||||
|
_creator_model_count_cache[cache_key] = (time.monotonic(), count)
|
||||||
|
return count
|
||||||
|
|||||||
@@ -201,6 +201,11 @@ PROVIDER_PRESETS: Dict[str, Dict[str, Any]] = {
|
|||||||
"api_base": "https://openrouter.ai/api/v1",
|
"api_base": "https://openrouter.ai/api/v1",
|
||||||
"requires_key": True,
|
"requires_key": True,
|
||||||
},
|
},
|
||||||
|
"google": {
|
||||||
|
"name": "Gemini",
|
||||||
|
"api_base": "https://generativelanguage.googleapis.com/v1beta/openai",
|
||||||
|
"requires_key": True,
|
||||||
|
},
|
||||||
"opencode-go": {
|
"opencode-go": {
|
||||||
"name": "OpenCode Go",
|
"name": "OpenCode Go",
|
||||||
"api_base": "https://opencode.ai/zen/go/v1",
|
"api_base": "https://opencode.ai/zen/go/v1",
|
||||||
|
|||||||
+21
-12
@@ -1,6 +1,7 @@
|
|||||||
import asyncio
|
import asyncio
|
||||||
import time
|
import time
|
||||||
import logging
|
import logging
|
||||||
|
import random
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
from typing import Any, Dict, List, Optional, Tuple
|
from typing import Any, Dict, List, Optional, Tuple
|
||||||
@@ -38,8 +39,8 @@ class ModelCache:
|
|||||||
|
|
||||||
def __post_init__(self):
|
def __post_init__(self):
|
||||||
self._lock = asyncio.Lock()
|
self._lock = asyncio.Lock()
|
||||||
# Cache for last sort: (sort_key, order) -> sorted list
|
# Cache for last sort: (sort_key, order, seed) -> sorted list
|
||||||
self._last_sort: Tuple[str, str] = (None, None)
|
self._last_sort: Tuple[Optional[str], str, Optional[str]] = (None, "asc", None)
|
||||||
self._last_sorted_data: List[Dict] = []
|
self._last_sorted_data: List[Dict] = []
|
||||||
self._normalize_raw_data()
|
self._normalize_raw_data()
|
||||||
self.name_display_mode = self._normalize_display_mode(self.name_display_mode)
|
self.name_display_mode = self._normalize_display_mode(self.name_display_mode)
|
||||||
@@ -203,9 +204,9 @@ class ModelCache:
|
|||||||
async def resort(self):
|
async def resort(self):
|
||||||
"""Resort cached data according to last sort mode if set"""
|
"""Resort cached data according to last sort mode if set"""
|
||||||
async with self._lock:
|
async with self._lock:
|
||||||
if self._last_sort != (None, None):
|
if self._last_sort[0] is not None:
|
||||||
sort_key, order = self._last_sort
|
sort_key, order, seed = self._last_sort
|
||||||
sorted_data = self._sort_data(self.raw_data, sort_key, order)
|
sorted_data = self._sort_data(self.raw_data, sort_key, order, seed)
|
||||||
self._last_sorted_data = sorted_data
|
self._last_sorted_data = sorted_data
|
||||||
# Update folder list
|
# Update folder list
|
||||||
# else: do nothing
|
# else: do nothing
|
||||||
@@ -218,7 +219,7 @@ class ModelCache:
|
|||||||
self.folders = sorted(list(all_folders), key=lambda x: x.lower())
|
self.folders = sorted(list(all_folders), key=lambda x: x.lower())
|
||||||
self.rebuild_version_index()
|
self.rebuild_version_index()
|
||||||
|
|
||||||
def _sort_data(self, data: List[Dict], sort_key: str, order: str) -> List[Dict]:
|
def _sort_data(self, data: List[Dict], sort_key: str, order: str, seed: Optional[str] = None) -> List[Dict]:
|
||||||
"""Sort data by sort_key and order"""
|
"""Sort data by sort_key and order"""
|
||||||
start_time = time.perf_counter()
|
start_time = time.perf_counter()
|
||||||
reverse = (order == 'desc')
|
reverse = (order == 'desc')
|
||||||
@@ -265,6 +266,13 @@ class ModelCache:
|
|||||||
),
|
),
|
||||||
reverse=reverse
|
reverse=reverse
|
||||||
)
|
)
|
||||||
|
elif sort_key == 'random':
|
||||||
|
# Random shuffle seeded for stable pagination: the same seed
|
||||||
|
# always yields the same order, so successive page requests
|
||||||
|
# stay consistent while browsing.
|
||||||
|
rng = random.Random(seed or 'random')
|
||||||
|
result = list(data)
|
||||||
|
rng.shuffle(result)
|
||||||
elif sort_key == 'versions_count':
|
elif sort_key == 'versions_count':
|
||||||
# Pre-dedup sort: fall back to name sort.
|
# Pre-dedup sort: fall back to name sort.
|
||||||
# Actual re-sort by version_count happens in get_paginated_data after dedup.
|
# Actual re-sort by version_count happens in get_paginated_data after dedup.
|
||||||
@@ -285,15 +293,16 @@ class ModelCache:
|
|||||||
logger.debug("ModelCache._sort_data(%s, %s) for %d items took %.3fs", sort_key, order, len(data), duration)
|
logger.debug("ModelCache._sort_data(%s, %s) for %d items took %.3fs", sort_key, order, len(data), duration)
|
||||||
return result
|
return result
|
||||||
|
|
||||||
async def get_sorted_data(self, sort_key: str = 'name', order: str = 'asc') -> List[Dict]:
|
async def get_sorted_data(self, sort_key: str = 'name', order: str = 'asc', seed: Optional[str] = None) -> List[Dict]:
|
||||||
"""Get sorted data by sort_key and order, using cache if possible"""
|
"""Get sorted data by sort_key and order, using cache if possible"""
|
||||||
async with self._lock:
|
async with self._lock:
|
||||||
if (sort_key, order) == self._last_sort:
|
cache_key = (sort_key, order, seed)
|
||||||
|
if cache_key == self._last_sort:
|
||||||
return self._last_sorted_data
|
return self._last_sorted_data
|
||||||
|
|
||||||
start_time = time.perf_counter()
|
start_time = time.perf_counter()
|
||||||
sorted_data = self._sort_data(self.raw_data, sort_key, order)
|
sorted_data = self._sort_data(self.raw_data, sort_key, order, seed)
|
||||||
self._last_sort = (sort_key, order)
|
self._last_sort = cache_key
|
||||||
self._last_sorted_data = sorted_data
|
self._last_sorted_data = sorted_data
|
||||||
|
|
||||||
duration = time.perf_counter() - start_time
|
duration = time.perf_counter() - start_time
|
||||||
@@ -313,8 +322,8 @@ class ModelCache:
|
|||||||
self.name_display_mode = normalized
|
self.name_display_mode = normalized
|
||||||
|
|
||||||
if self._last_sort[0] == 'name':
|
if self._last_sort[0] == 'name':
|
||||||
sort_key, order = self._last_sort
|
sort_key, order, seed = self._last_sort
|
||||||
self._last_sorted_data = self._sort_data(self.raw_data, sort_key, order)
|
self._last_sorted_data = self._sort_data(self.raw_data, sort_key, order, seed)
|
||||||
|
|
||||||
async def update_preview_url(self, file_path: str, preview_url: str, preview_nsfw_level: int) -> bool:
|
async def update_preview_url(self, file_path: str, preview_url: str, preview_nsfw_level: int) -> bool:
|
||||||
"""Update preview_url for a specific model in all cached data
|
"""Update preview_url for a specific model in all cached data
|
||||||
|
|||||||
@@ -143,10 +143,18 @@ class ModelMetadataProvider(ABC):
|
|||||||
pass
|
pass
|
||||||
|
|
||||||
@abstractmethod
|
@abstractmethod
|
||||||
async def get_user_models(self, username: str) -> Optional[List[Dict]]:
|
async def get_user_models(self, username: str, cursor: Optional[str] = None) -> Optional[Dict]:
|
||||||
"""Fetch models owned by the specified user"""
|
"""Fetch one page of models owned by the specified user.
|
||||||
|
|
||||||
|
Returns ``{"items": [...], "nextCursor": <str|None>}`` on success,
|
||||||
|
or None when unsupported/failed. ``cursor`` continues a previous page.
|
||||||
|
"""
|
||||||
pass
|
pass
|
||||||
|
|
||||||
|
async def get_creator_model_count(self, username: str) -> Optional[int]:
|
||||||
|
"""Published model count for the user; None when unsupported."""
|
||||||
|
return None
|
||||||
|
|
||||||
class CivitaiModelMetadataProvider(ModelMetadataProvider):
|
class CivitaiModelMetadataProvider(ModelMetadataProvider):
|
||||||
"""Provider that uses Civitai API for metadata"""
|
"""Provider that uses Civitai API for metadata"""
|
||||||
|
|
||||||
@@ -175,8 +183,11 @@ class CivitaiModelMetadataProvider(ModelMetadataProvider):
|
|||||||
async def get_model_version_info(self, version_id: str) -> Tuple[Optional[Dict], Optional[str]]:
|
async def get_model_version_info(self, version_id: str) -> Tuple[Optional[Dict], Optional[str]]:
|
||||||
return await self.client.get_model_version_info(version_id)
|
return await self.client.get_model_version_info(version_id)
|
||||||
|
|
||||||
async def get_user_models(self, username: str) -> Optional[List[Dict]]:
|
async def get_user_models(self, username: str, cursor: Optional[str] = None) -> Optional[Dict]:
|
||||||
return await self.client.get_user_models(username)
|
return await self.client.get_user_models(username, cursor)
|
||||||
|
|
||||||
|
async def get_creator_model_count(self, username: str) -> Optional[int]:
|
||||||
|
return await self.client.get_creator_model_count(username)
|
||||||
|
|
||||||
class CivArchiveModelMetadataProvider(ModelMetadataProvider):
|
class CivArchiveModelMetadataProvider(ModelMetadataProvider):
|
||||||
"""Provider that uses CivArchive API for metadata"""
|
"""Provider that uses CivArchive API for metadata"""
|
||||||
@@ -196,7 +207,7 @@ class CivArchiveModelMetadataProvider(ModelMetadataProvider):
|
|||||||
async def get_model_version_info(self, version_id: str) -> Tuple[Optional[Dict], Optional[str]]:
|
async def get_model_version_info(self, version_id: str) -> Tuple[Optional[Dict], Optional[str]]:
|
||||||
return await self.client.get_model_version_info(version_id)
|
return await self.client.get_model_version_info(version_id)
|
||||||
|
|
||||||
async def get_user_models(self, username: str) -> Optional[List[Dict]]:
|
async def get_user_models(self, username: str, cursor: Optional[str] = None) -> Optional[Dict]:
|
||||||
"""Not supported by CivArchive provider"""
|
"""Not supported by CivArchive provider"""
|
||||||
return None
|
return None
|
||||||
|
|
||||||
@@ -347,7 +358,7 @@ class SQLiteModelMetadataProvider(ModelMetadataProvider):
|
|||||||
version_data = await self._get_version_with_model_data(db, model_id, version_id)
|
version_data = await self._get_version_with_model_data(db, model_id, version_id)
|
||||||
return version_data, None
|
return version_data, None
|
||||||
|
|
||||||
async def get_user_models(self, username: str) -> Optional[List[Dict]]:
|
async def get_user_models(self, username: str, cursor: Optional[str] = None) -> Optional[Dict]:
|
||||||
"""Listing models by username is not supported for archive database"""
|
"""Listing models by username is not supported for archive database"""
|
||||||
return None
|
return None
|
||||||
|
|
||||||
@@ -602,13 +613,14 @@ class FallbackMetadataProvider(ModelMetadataProvider):
|
|||||||
continue
|
continue
|
||||||
return None
|
return None
|
||||||
|
|
||||||
async def get_user_models(self, username: str) -> Optional[List[Dict]]:
|
async def get_user_models(self, username: str, cursor: Optional[str] = None) -> Optional[Dict]:
|
||||||
for provider, label in self._iter_providers():
|
for provider, label in self._iter_providers():
|
||||||
try:
|
try:
|
||||||
result = await self._call_with_rate_limit(
|
result = await self._call_with_rate_limit(
|
||||||
label,
|
label,
|
||||||
provider.get_user_models,
|
provider.get_user_models,
|
||||||
username,
|
username,
|
||||||
|
cursor=cursor,
|
||||||
)
|
)
|
||||||
if result is not None:
|
if result is not None:
|
||||||
return result
|
return result
|
||||||
@@ -624,6 +636,19 @@ class FallbackMetadataProvider(ModelMetadataProvider):
|
|||||||
continue
|
continue
|
||||||
return None
|
return None
|
||||||
|
|
||||||
|
async def get_creator_model_count(self, username: str) -> Optional[int]:
|
||||||
|
for provider, label in self._iter_providers():
|
||||||
|
try:
|
||||||
|
result = await provider.get_creator_model_count(username)
|
||||||
|
if result is not None:
|
||||||
|
return result
|
||||||
|
except Exception as e:
|
||||||
|
logger.debug(
|
||||||
|
"Provider %s failed for get_creator_model_count: %s", label, e
|
||||||
|
)
|
||||||
|
continue
|
||||||
|
return None
|
||||||
|
|
||||||
def _iter_providers(self):
|
def _iter_providers(self):
|
||||||
return zip(self.providers, self._provider_labels)
|
return zip(self.providers, self._provider_labels)
|
||||||
|
|
||||||
@@ -704,13 +729,17 @@ class RateLimitRetryingProvider(ModelMetadataProvider):
|
|||||||
version_id,
|
version_id,
|
||||||
)
|
)
|
||||||
|
|
||||||
async def get_user_models(self, username: str) -> Optional[List[Dict]]:
|
async def get_user_models(self, username: str, cursor: Optional[str] = None) -> Optional[Dict]:
|
||||||
return await self._rate_limit_helper.run(
|
return await self._rate_limit_helper.run(
|
||||||
self._label,
|
self._label,
|
||||||
self._provider.get_user_models,
|
self._provider.get_user_models,
|
||||||
username,
|
username,
|
||||||
|
cursor=cursor,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
async def get_creator_model_count(self, username: str) -> Optional[int]:
|
||||||
|
return await self._provider.get_creator_model_count(username)
|
||||||
|
|
||||||
class ModelMetadataProviderManager:
|
class ModelMetadataProviderManager:
|
||||||
"""Manager for selecting and using model metadata providers"""
|
"""Manager for selecting and using model metadata providers"""
|
||||||
|
|
||||||
@@ -776,10 +805,20 @@ class ModelMetadataProviderManager:
|
|||||||
except NotImplementedError:
|
except NotImplementedError:
|
||||||
return None
|
return None
|
||||||
|
|
||||||
async def get_user_models(self, username: str, provider_name: str = None) -> Optional[List[Dict]]:
|
async def get_user_models(
|
||||||
"""Fetch models owned by the specified user"""
|
self,
|
||||||
|
username: str,
|
||||||
|
provider_name: str = None,
|
||||||
|
cursor: Optional[str] = None,
|
||||||
|
) -> Optional[Dict]:
|
||||||
|
"""Fetch one page of models owned by the specified user"""
|
||||||
provider = self._get_provider(provider_name)
|
provider = self._get_provider(provider_name)
|
||||||
return await provider.get_user_models(username)
|
return await provider.get_user_models(username, cursor)
|
||||||
|
|
||||||
|
async def get_creator_model_count(self, username: str, provider_name: str = None) -> Optional[int]:
|
||||||
|
"""Best-effort published model count for the specified user"""
|
||||||
|
provider = self._get_provider(provider_name)
|
||||||
|
return await provider.get_creator_model_count(username)
|
||||||
|
|
||||||
def _get_provider(self, provider_name: str = None) -> ModelMetadataProvider:
|
def _get_provider(self, provider_name: str = None) -> ModelMetadataProvider:
|
||||||
"""Get provider by name or default provider"""
|
"""Get provider by name or default provider"""
|
||||||
|
|||||||
@@ -85,6 +85,7 @@ class SortParams:
|
|||||||
|
|
||||||
key: str
|
key: str
|
||||||
order: str
|
order: str
|
||||||
|
seed: Optional[str] = None
|
||||||
|
|
||||||
|
|
||||||
@dataclass(frozen=True)
|
@dataclass(frozen=True)
|
||||||
@@ -116,7 +117,7 @@ class ModelCacheRepository:
|
|||||||
async def fetch_sorted(self, params: SortParams) -> List[Dict[str, Any]]:
|
async def fetch_sorted(self, params: SortParams) -> List[Dict[str, Any]]:
|
||||||
"""Fetch cached data pre-sorted according to ``params``."""
|
"""Fetch cached data pre-sorted according to ``params``."""
|
||||||
cache = await self.get_cache()
|
cache = await self.get_cache()
|
||||||
return await cache.get_sorted_data(params.key, params.order)
|
return await cache.get_sorted_data(params.key, params.order, params.seed)
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def parse_sort(sort_by: str) -> SortParams:
|
def parse_sort(sort_by: str) -> SortParams:
|
||||||
@@ -132,10 +133,17 @@ class ModelCacheRepository:
|
|||||||
sort_key = sort_by.strip().lower() or "name"
|
sort_key = sort_by.strip().lower() or "name"
|
||||||
order = "asc"
|
order = "asc"
|
||||||
|
|
||||||
if order not in ("asc", "desc"):
|
seed = None
|
||||||
|
if sort_key == "random":
|
||||||
|
# Random sort: the portion after ':' is the shuffle seed.
|
||||||
|
# A stable seed keeps paginated requests consistent; order is
|
||||||
|
# meaningless for a random shuffle.
|
||||||
|
seed = order if order and order not in ("asc", "desc") else None
|
||||||
|
order = "asc"
|
||||||
|
elif order not in ("asc", "desc"):
|
||||||
order = "asc"
|
order = "asc"
|
||||||
|
|
||||||
return SortParams(key=sort_key, order=order)
|
return SortParams(key=sort_key, order=order, seed=seed)
|
||||||
|
|
||||||
|
|
||||||
class ModelFilterSet:
|
class ModelFilterSet:
|
||||||
|
|||||||
@@ -1752,7 +1752,7 @@ class ModelScanner:
|
|||||||
# ---- Conditional resort (only when sort-key fields changed) ----
|
# ---- Conditional resort (only when sort-key fields changed) ----
|
||||||
need_resort = False
|
need_resort = False
|
||||||
_last = cache._last_sort
|
_last = cache._last_sort
|
||||||
sort_key: Optional[str] = _last[0] if _last != (None, None) else None
|
sort_key: Optional[str] = _last[0] if _last[0] is not None else None
|
||||||
if sort_key == "name":
|
if sort_key == "name":
|
||||||
if (
|
if (
|
||||||
old_model_name != desired_entry.get("model_name", "")
|
old_model_name != desired_entry.get("model_name", "")
|
||||||
|
|||||||
@@ -92,6 +92,7 @@ DEFAULT_SETTINGS: Dict[str, Any] = {
|
|||||||
"mature_blur_level": "R",
|
"mature_blur_level": "R",
|
||||||
"autoplay_on_hover": False,
|
"autoplay_on_hover": False,
|
||||||
"display_density": "default",
|
"display_density": "default",
|
||||||
|
"recipes_layout": "grid",
|
||||||
"card_info_display": "always",
|
"card_info_display": "always",
|
||||||
"include_trigger_words": False,
|
"include_trigger_words": False,
|
||||||
"compact_mode": False,
|
"compact_mode": False,
|
||||||
@@ -1473,10 +1474,12 @@ class SettingsManager:
|
|||||||
|
|
||||||
try:
|
try:
|
||||||
common_root = os.path.commonpath([source, target])
|
common_root = os.path.commonpath([source, target])
|
||||||
except ValueError as exc:
|
except ValueError:
|
||||||
raise ValueError("Invalid recipes path change") from exc
|
# Windows: paths on different drives share no common root.
|
||||||
|
# A cross-drive move is valid, so treat it as no common root.
|
||||||
|
common_root = None
|
||||||
|
|
||||||
if common_root == source:
|
if common_root is not None and common_root == source:
|
||||||
raise ValueError("Recipes path cannot be moved into a nested directory")
|
raise ValueError("Recipes path cannot be moved into a nested directory")
|
||||||
|
|
||||||
planned_recipe_updates: Dict[str, Dict[str, Any]] = {}
|
planned_recipe_updates: Dict[str, Dict[str, Any]] = {}
|
||||||
|
|||||||
@@ -14,11 +14,16 @@ from ..services.service_registry import ServiceRegistry
|
|||||||
from ..utils.example_images_paths import (
|
from ..utils.example_images_paths import (
|
||||||
ExampleImagePathResolver,
|
ExampleImagePathResolver,
|
||||||
ensure_library_root_exists,
|
ensure_library_root_exists,
|
||||||
|
get_example_images_root,
|
||||||
|
is_hash_folder,
|
||||||
uses_library_scoped_folders,
|
uses_library_scoped_folders,
|
||||||
)
|
)
|
||||||
from ..utils.metadata_manager import MetadataManager
|
from ..utils.metadata_manager import MetadataManager
|
||||||
from .example_images_processor import ExampleImagesProcessor
|
from .example_images_processor import ExampleImagesProcessor
|
||||||
from .example_images_metadata import MetadataUpdater
|
from .example_images_metadata import (
|
||||||
|
MetadataUpdater,
|
||||||
|
update_cache_from_metadata,
|
||||||
|
)
|
||||||
from ..services.downloader import get_downloader
|
from ..services.downloader import get_downloader
|
||||||
from ..services.settings_manager import get_settings_manager
|
from ..services.settings_manager import get_settings_manager
|
||||||
|
|
||||||
@@ -87,6 +92,13 @@ class _DownloadProgress(dict):
|
|||||||
return snapshot
|
return snapshot
|
||||||
|
|
||||||
|
|
||||||
|
# When fewer candidates than this remain in check_pending_models, probe each
|
||||||
|
# model folder directly (preserving legacy-folder migration semantics). Above
|
||||||
|
# it, build a folder index with a single directory scan so libraries with
|
||||||
|
# 100k+ models do not pay one syscall per candidate.
|
||||||
|
_BULK_LOOKUP_THRESHOLD = 1000
|
||||||
|
|
||||||
|
|
||||||
def _model_directory_has_files(path: str) -> bool:
|
def _model_directory_has_files(path: str) -> bool:
|
||||||
"""Return True when the provided directory exists and contains entries."""
|
"""Return True when the provided directory exists and contains entries."""
|
||||||
|
|
||||||
@@ -103,6 +115,36 @@ def _model_directory_has_files(path: str) -> bool:
|
|||||||
return False
|
return False
|
||||||
|
|
||||||
|
|
||||||
|
def _build_example_folder_index(output_dir: str) -> dict[str, bool]:
|
||||||
|
"""Build a ``{hash: has_files}`` index for a library's example-image folders.
|
||||||
|
|
||||||
|
A single directory scan over the library root replaces ``O(candidates)``
|
||||||
|
per-folder ``os.scandir`` calls, which is required for libraries with
|
||||||
|
100k+ models. Each hash folder is classified by whether it contains any
|
||||||
|
entries, matching the semantics of ``_model_directory_has_files``.
|
||||||
|
"""
|
||||||
|
|
||||||
|
index: dict[str, bool] = {}
|
||||||
|
if not output_dir or not os.path.isdir(output_dir):
|
||||||
|
return index
|
||||||
|
|
||||||
|
try:
|
||||||
|
with os.scandir(output_dir) as entries:
|
||||||
|
for entry in entries:
|
||||||
|
name = entry.name
|
||||||
|
if not entry.is_dir() or not is_hash_folder(name):
|
||||||
|
continue
|
||||||
|
try:
|
||||||
|
with os.scandir(entry.path) as subentries:
|
||||||
|
index[name.lower()] = any(subentries)
|
||||||
|
except OSError:
|
||||||
|
index[name.lower()] = False
|
||||||
|
except OSError:
|
||||||
|
pass
|
||||||
|
|
||||||
|
return index
|
||||||
|
|
||||||
|
|
||||||
class DownloadManager:
|
class DownloadManager:
|
||||||
"""Manages downloading example images for models."""
|
"""Manages downloading example images for models."""
|
||||||
|
|
||||||
@@ -130,6 +172,7 @@ class DownloadManager:
|
|||||||
model_types = data.get("model_types", ["lora", "checkpoint"])
|
model_types = data.get("model_types", ["lora", "checkpoint"])
|
||||||
delay = float(data.get("delay", 0.2))
|
delay = float(data.get("delay", 0.2))
|
||||||
force = data.get("force", False)
|
force = data.get("force", False)
|
||||||
|
model_hashes = data.get("model_hashes", [])
|
||||||
|
|
||||||
# Step 2: Validate configuration (fast lookup)
|
# Step 2: Validate configuration (fast lookup)
|
||||||
settings_manager = get_settings_manager()
|
settings_manager = get_settings_manager()
|
||||||
@@ -199,6 +242,7 @@ class DownloadManager:
|
|||||||
delay,
|
delay,
|
||||||
active_library,
|
active_library,
|
||||||
force,
|
force,
|
||||||
|
model_hashes,
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -410,14 +454,49 @@ class DownloadManager:
|
|||||||
# Calculate pending count: check which models actually need processing.
|
# Calculate pending count: check which models actually need processing.
|
||||||
# A model is pending if it has a hash, is not already processed or known-failed,
|
# A model is pending if it has a hash, is not already processed or known-failed,
|
||||||
# and its folder doesn't exist or is empty.
|
# and its folder doesn't exist or is empty.
|
||||||
pending_hashes = set()
|
candidate_hashes = [
|
||||||
for model_hash, model_name in all_models_with_hash:
|
model_hash
|
||||||
if model_hash not in processed_models and model_hash not in failed_models:
|
for model_hash, _ in all_models_with_hash
|
||||||
|
if model_hash not in processed_models
|
||||||
|
and model_hash not in failed_models
|
||||||
|
]
|
||||||
|
|
||||||
|
pending_hashes: set[str] = set()
|
||||||
|
# For small candidate counts the existing per-folder check is fine
|
||||||
|
# and handles legacy folder migration.
|
||||||
|
# For large libraries, scan the library root once and do set lookups.
|
||||||
|
if len(candidate_hashes) <= _BULK_LOOKUP_THRESHOLD or not output_dir:
|
||||||
|
for model_hash in candidate_hashes:
|
||||||
model_dir = ExampleImagePathResolver.get_model_folder(
|
model_dir = ExampleImagePathResolver.get_model_folder(
|
||||||
model_hash, active_library
|
model_hash, active_library
|
||||||
)
|
)
|
||||||
if not _model_directory_has_files(model_dir):
|
if not _model_directory_has_files(model_dir):
|
||||||
pending_hashes.add(model_hash)
|
pending_hashes.add(model_hash)
|
||||||
|
else:
|
||||||
|
folder_index = await asyncio.get_event_loop().run_in_executor(
|
||||||
|
None, _build_example_folder_index, output_dir
|
||||||
|
)
|
||||||
|
# In multi-library mode, folders that have not been consolidated
|
||||||
|
# into the library root yet (startup migration skipped, failed
|
||||||
|
# move, or created at the legacy path afterwards) still live at
|
||||||
|
# the legacy root/<hash> location. Only scan that root when at
|
||||||
|
# least one candidate is missing from the library-root index, so
|
||||||
|
# the fully-consolidated case does not pay an extra directory
|
||||||
|
# pass on every call.
|
||||||
|
if uses_library_scoped_folders() and any(
|
||||||
|
not folder_index.get(model_hash, False)
|
||||||
|
for model_hash in candidate_hashes
|
||||||
|
):
|
||||||
|
legacy_root = get_example_images_root()
|
||||||
|
if legacy_root and legacy_root != output_dir:
|
||||||
|
legacy_index = await asyncio.get_event_loop().run_in_executor(
|
||||||
|
None, _build_example_folder_index, legacy_root
|
||||||
|
)
|
||||||
|
for hash_key, has_files in legacy_index.items():
|
||||||
|
folder_index.setdefault(hash_key, has_files)
|
||||||
|
for model_hash in candidate_hashes:
|
||||||
|
if not folder_index.get(model_hash, False):
|
||||||
|
pending_hashes.add(model_hash)
|
||||||
|
|
||||||
pending_count = len(pending_hashes)
|
pending_count = len(pending_hashes)
|
||||||
|
|
||||||
@@ -500,8 +579,9 @@ class DownloadManager:
|
|||||||
delay,
|
delay,
|
||||||
library_name,
|
library_name,
|
||||||
force: bool = False,
|
force: bool = False,
|
||||||
|
model_hashes: list[str] | None = None,
|
||||||
):
|
):
|
||||||
"""Download example images for all models."""
|
"""Download example images for all models (or only the given hashes)."""
|
||||||
|
|
||||||
downloader = await get_downloader()
|
downloader = await get_downloader()
|
||||||
|
|
||||||
@@ -529,6 +609,18 @@ class DownloadManager:
|
|||||||
if model.get("sha256"):
|
if model.get("sha256"):
|
||||||
all_models.append((scanner_type, model, scanner))
|
all_models.append((scanner_type, model, scanner))
|
||||||
|
|
||||||
|
# Restrict to the requested hashes when provided (empty = all models).
|
||||||
|
# Explicit targets are a directed user request, so previously failed
|
||||||
|
# models are retried instead of skipped.
|
||||||
|
explicit_targets = bool(model_hashes)
|
||||||
|
if model_hashes:
|
||||||
|
hash_set = {h.lower() for h in model_hashes}
|
||||||
|
all_models = [
|
||||||
|
(scanner_type, model, scanner)
|
||||||
|
for scanner_type, model, scanner in all_models
|
||||||
|
if model.get("sha256", "").lower() in hash_set
|
||||||
|
]
|
||||||
|
|
||||||
# Update total count
|
# Update total count
|
||||||
self._progress["total"] = len(all_models)
|
self._progress["total"] = len(all_models)
|
||||||
logger.debug(f"Found {self._progress['total']} models to process")
|
logger.debug(f"Found {self._progress['total']} models to process")
|
||||||
@@ -552,6 +644,7 @@ class DownloadManager:
|
|||||||
downloader,
|
downloader,
|
||||||
library_name,
|
library_name,
|
||||||
force,
|
force,
|
||||||
|
explicit_targets,
|
||||||
)
|
)
|
||||||
|
|
||||||
# Update progress
|
# Update progress
|
||||||
@@ -648,6 +741,7 @@ class DownloadManager:
|
|||||||
downloader,
|
downloader,
|
||||||
library_name,
|
library_name,
|
||||||
force: bool = False,
|
force: bool = False,
|
||||||
|
explicit_targets: bool = False,
|
||||||
):
|
):
|
||||||
"""Process a single model download."""
|
"""Process a single model download."""
|
||||||
|
|
||||||
@@ -670,8 +764,9 @@ class DownloadManager:
|
|||||||
self._progress["current_model"] = f"{model_name} ({model_hash[:8]})"
|
self._progress["current_model"] = f"{model_name} ({model_hash[:8]})"
|
||||||
await self._broadcast_progress(status="running")
|
await self._broadcast_progress(status="running")
|
||||||
|
|
||||||
# Skip if already in failed models (unless force mode is enabled)
|
# Skip if already in failed models (unless force mode is enabled or
|
||||||
if not force and model_hash in self._progress["failed_models"]:
|
# the model was explicitly targeted by hash)
|
||||||
|
if not force and not explicit_targets and model_hash in self._progress["failed_models"]:
|
||||||
logger.debug(f"Skipping known failed model: {model_name}")
|
logger.debug(f"Skipping known failed model: {model_name}")
|
||||||
return False
|
return False
|
||||||
|
|
||||||
@@ -680,30 +775,34 @@ class DownloadManager:
|
|||||||
)
|
)
|
||||||
existing_files = _model_directory_has_files(model_dir)
|
existing_files = _model_directory_has_files(model_dir)
|
||||||
|
|
||||||
# Skip if already processed AND directory exists with files
|
# Model-level guard: a populated folder counts as done. Explicitly
|
||||||
if model_hash in self._progress["processed_models"]:
|
# targeted models bypass it so the per-image existence pre-check can
|
||||||
if existing_files:
|
# fill individual gaps without re-fetching existing files.
|
||||||
logger.debug(f"Skipping already processed model: {model_name}")
|
if not explicit_targets:
|
||||||
|
# Skip if already processed AND directory exists with files
|
||||||
|
if model_hash in self._progress["processed_models"]:
|
||||||
|
if existing_files:
|
||||||
|
logger.debug(f"Skipping already processed model: {model_name}")
|
||||||
|
return False
|
||||||
|
|
||||||
|
logger.debug(
|
||||||
|
"Model %s (%s) marked as processed but folder empty or missing, reprocessing triggered",
|
||||||
|
model_name,
|
||||||
|
model_hash,
|
||||||
|
)
|
||||||
|
# Track that we are reprocessing this model for summary logging
|
||||||
|
self._progress["reprocessed_models"].add(model_hash)
|
||||||
|
# Remove from processed models since we need to reprocess
|
||||||
|
self._progress["processed_models"].discard(model_hash)
|
||||||
|
|
||||||
|
if existing_files and model_hash not in self._progress["processed_models"]:
|
||||||
|
logger.debug(
|
||||||
|
"Model folder already populated for %s, marking as processed without download",
|
||||||
|
model_name,
|
||||||
|
)
|
||||||
|
self._progress["processed_models"].add(model_hash)
|
||||||
return False
|
return False
|
||||||
|
|
||||||
logger.debug(
|
|
||||||
"Model %s (%s) marked as processed but folder empty or missing, reprocessing triggered",
|
|
||||||
model_name,
|
|
||||||
model_hash,
|
|
||||||
)
|
|
||||||
# Track that we are reprocessing this model for summary logging
|
|
||||||
self._progress["reprocessed_models"].add(model_hash)
|
|
||||||
# Remove from processed models since we need to reprocess
|
|
||||||
self._progress["processed_models"].discard(model_hash)
|
|
||||||
|
|
||||||
if existing_files and model_hash not in self._progress["processed_models"]:
|
|
||||||
logger.debug(
|
|
||||||
"Model folder already populated for %s, marking as processed without download",
|
|
||||||
model_name,
|
|
||||||
)
|
|
||||||
self._progress["processed_models"].add(model_hash)
|
|
||||||
return False
|
|
||||||
|
|
||||||
if not model_dir:
|
if not model_dir:
|
||||||
logger.warning(
|
logger.warning(
|
||||||
"Unable to resolve example images folder for model %s (%s)",
|
"Unable to resolve example images folder for model %s (%s)",
|
||||||
@@ -807,7 +906,7 @@ class DownloadManager:
|
|||||||
model_name,
|
model_name,
|
||||||
)
|
)
|
||||||
# Clear failed_models so non-force runs can retry
|
# Clear failed_models so non-force runs can retry
|
||||||
if force and model_hash in self._progress["failed_models"]:
|
if (force or explicit_targets) and model_hash in self._progress["failed_models"]:
|
||||||
self._progress["failed_models"].discard(model_hash)
|
self._progress["failed_models"].discard(model_hash)
|
||||||
logger.info(
|
logger.info(
|
||||||
f"Removed {model_name} from failed_models after force retry with rate-limited images"
|
f"Removed {model_name} from failed_models after force retry with rate-limited images"
|
||||||
@@ -827,7 +926,7 @@ class DownloadManager:
|
|||||||
)
|
)
|
||||||
elif success:
|
elif success:
|
||||||
self._progress["processed_models"].add(model_hash)
|
self._progress["processed_models"].add(model_hash)
|
||||||
if force and model_hash in self._progress["failed_models"]:
|
if (force or explicit_targets) and model_hash in self._progress["failed_models"]:
|
||||||
self._progress["failed_models"].discard(model_hash)
|
self._progress["failed_models"].discard(model_hash)
|
||||||
logger.info(
|
logger.info(
|
||||||
f"Removed {model_name} from failed_models after successful force retry"
|
f"Removed {model_name} from failed_models after successful force retry"
|
||||||
@@ -1343,8 +1442,8 @@ class DownloadManager:
|
|||||||
await MetadataManager.save_metadata(file_path, model_copy)
|
await MetadataManager.save_metadata(file_path, model_copy)
|
||||||
|
|
||||||
try:
|
try:
|
||||||
await scanner.update_single_model_cache(
|
await update_cache_from_metadata(
|
||||||
file_path, file_path, model_data
|
scanner, file_path, model_copy
|
||||||
)
|
)
|
||||||
except AttributeError:
|
except AttributeError:
|
||||||
logger.debug(
|
logger.debug(
|
||||||
|
|||||||
@@ -1,3 +1,4 @@
|
|||||||
|
import inspect
|
||||||
import logging
|
import logging
|
||||||
import os
|
import os
|
||||||
import re
|
import re
|
||||||
@@ -28,6 +29,31 @@ if TYPE_CHECKING: # pragma: no cover - import for type checkers only
|
|||||||
from ..services.settings_manager import SettingsManager
|
from ..services.settings_manager import SettingsManager
|
||||||
|
|
||||||
|
|
||||||
|
async def update_cache_from_metadata(
|
||||||
|
scanner: Any, file_path: str, metadata: Dict[str, Any]
|
||||||
|
) -> bool:
|
||||||
|
"""Update the scanner cache from a metadata dict using the in-place sync path.
|
||||||
|
|
||||||
|
``sync_cache_from_metadata`` patches the existing cache entry incrementally
|
||||||
|
(tag/hash/version indexes, targeted single-row SQL update) and only resorts
|
||||||
|
when a sort-key field changed. This avoids the ``O(n)`` full-list resort and
|
||||||
|
full cache rewrite that ``update_single_model_cache`` performs on every call,
|
||||||
|
which is critical for libraries with 100k+ models.
|
||||||
|
|
||||||
|
Falls back to the legacy full update when the scanner does not expose an
|
||||||
|
async ``sync_cache_from_metadata`` method.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
``True`` if the cache entry was updated, ``False`` otherwise.
|
||||||
|
"""
|
||||||
|
|
||||||
|
sync_method = getattr(scanner, "sync_cache_from_metadata", None)
|
||||||
|
if inspect.iscoroutinefunction(sync_method):
|
||||||
|
return await sync_method(file_path, metadata)
|
||||||
|
|
||||||
|
return await scanner.update_single_model_cache(file_path, file_path, metadata)
|
||||||
|
|
||||||
|
|
||||||
def _build_metadata_sync_service(settings_manager: "SettingsManager") -> MetadataSyncService:
|
def _build_metadata_sync_service(settings_manager: "SettingsManager") -> MetadataSyncService:
|
||||||
"""Construct a metadata sync service bound to the provided settings."""
|
"""Construct a metadata sync service bound to the provided settings."""
|
||||||
|
|
||||||
@@ -103,8 +129,8 @@ class MetadataUpdater:
|
|||||||
progress['refreshed_models'].add(model_hash)
|
progress['refreshed_models'].add(model_hash)
|
||||||
|
|
||||||
async def update_cache_func(old_path, new_path, metadata):
|
async def update_cache_func(old_path, new_path, metadata):
|
||||||
return await scanner.update_single_model_cache(old_path, new_path, metadata)
|
return await update_cache_from_metadata(scanner, new_path, metadata)
|
||||||
|
|
||||||
await MetadataManager.hydrate_model_data(model_data)
|
await MetadataManager.hydrate_model_data(model_data)
|
||||||
success, error = await _get_metadata_sync_service().fetch_and_update_model(
|
success, error = await _get_metadata_sync_service().fetch_and_update_model(
|
||||||
sha256=model_hash,
|
sha256=model_hash,
|
||||||
@@ -234,6 +260,7 @@ class MetadataUpdater:
|
|||||||
|
|
||||||
# Save metadata to .metadata.json file
|
# Save metadata to .metadata.json file
|
||||||
file_path = model.get('file_path')
|
file_path = model.get('file_path')
|
||||||
|
model_copy: Optional[Dict[str, Any]] = None
|
||||||
try:
|
try:
|
||||||
model_copy = model.copy()
|
model_copy = model.copy()
|
||||||
model_copy.pop('folder', None)
|
model_copy.pop('folder', None)
|
||||||
@@ -241,14 +268,18 @@ class MetadataUpdater:
|
|||||||
logger.info(f"Saved metadata for {model.get('model_name')}")
|
logger.info(f"Saved metadata for {model.get('model_name')}")
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.error(f"Failed to save metadata for {model.get('model_name')}: {str(e)}")
|
logger.error(f"Failed to save metadata for {model.get('model_name')}: {str(e)}")
|
||||||
|
|
||||||
# Save updated metadata to scanner cache
|
# Save updated metadata to scanner cache. sync_cache_from_metadata
|
||||||
success = await scanner.update_single_model_cache(file_path, file_path, model)
|
# returns False both for "already in sync" and for actual failures,
|
||||||
if success:
|
# so the cache sync result is deliberately not treated as an error;
|
||||||
|
# the return value reflects whether the metadata was persisted.
|
||||||
|
if file_path and model_copy is not None:
|
||||||
|
await update_cache_from_metadata(scanner, file_path, model_copy)
|
||||||
logger.info(f"Successfully updated metadata for {model.get('model_name')} with {len(images)} local examples")
|
logger.info(f"Successfully updated metadata for {model.get('model_name')} with {len(images)} local examples")
|
||||||
return True
|
return True
|
||||||
else:
|
|
||||||
logger.warning(f"Failed to update metadata for {model.get('model_name')}")
|
logger.warning(f"Failed to update metadata for {model.get('model_name')}")
|
||||||
|
return False
|
||||||
|
|
||||||
return False
|
return False
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
@@ -336,6 +367,7 @@ class MetadataUpdater:
|
|||||||
|
|
||||||
# Save metadata to .metadata.json file
|
# Save metadata to .metadata.json file
|
||||||
file_path = model_data.get('file_path')
|
file_path = model_data.get('file_path')
|
||||||
|
model_copy: Optional[Dict[str, Any]] = None
|
||||||
if file_path:
|
if file_path:
|
||||||
try:
|
try:
|
||||||
model_copy = model_data.copy()
|
model_copy = model_data.copy()
|
||||||
@@ -344,11 +376,11 @@ class MetadataUpdater:
|
|||||||
logger.info(f"Saved metadata for {model_data.get('model_name')}")
|
logger.info(f"Saved metadata for {model_data.get('model_name')}")
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.error(f"Failed to save metadata: {str(e)}")
|
logger.error(f"Failed to save metadata: {str(e)}")
|
||||||
|
|
||||||
# Save updated metadata to scanner cache
|
# Save updated metadata to scanner cache
|
||||||
if file_path:
|
if file_path and model_copy is not None:
|
||||||
await scanner.update_single_model_cache(file_path, file_path, model_data)
|
await update_cache_from_metadata(scanner, file_path, model_copy)
|
||||||
|
|
||||||
# Get regular images array (might be None)
|
# Get regular images array (might be None)
|
||||||
regular_images = civitai_data.get('images', [])
|
regular_images = civitai_data.get('images', [])
|
||||||
|
|
||||||
|
|||||||
@@ -15,6 +15,7 @@ from ..utils.example_images_paths import (
|
|||||||
)
|
)
|
||||||
from ..utils.metadata_manager import MetadataManager
|
from ..utils.metadata_manager import MetadataManager
|
||||||
from ..utils.example_images_processor import ExampleImagesProcessor
|
from ..utils.example_images_processor import ExampleImagesProcessor
|
||||||
|
from ..utils.example_images_metadata import update_cache_from_metadata
|
||||||
from ..utils.constants import SUPPORTED_MEDIA_EXTENSIONS
|
from ..utils.constants import SUPPORTED_MEDIA_EXTENSIONS
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
@@ -421,7 +422,7 @@ class ExampleImagesMigration:
|
|||||||
await MetadataManager.save_metadata(file_path, model_copy)
|
await MetadataManager.save_metadata(file_path, model_copy)
|
||||||
|
|
||||||
# Update scanner cache
|
# Update scanner cache
|
||||||
await scanner.update_single_model_cache(file_path, file_path, model_metadata)
|
await update_cache_from_metadata(scanner, file_path, model_copy)
|
||||||
|
|
||||||
updated_models += 1
|
updated_models += 1
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
|
|||||||
@@ -9,7 +9,7 @@ from ..utils.constants import SUPPORTED_MEDIA_EXTENSIONS
|
|||||||
from ..services.service_registry import ServiceRegistry
|
from ..services.service_registry import ServiceRegistry
|
||||||
from ..services.settings_manager import get_settings_manager
|
from ..services.settings_manager import get_settings_manager
|
||||||
from ..utils.example_images_paths import get_model_folder, get_model_relative_path
|
from ..utils.example_images_paths import get_model_folder, get_model_relative_path
|
||||||
from .example_images_metadata import MetadataUpdater
|
from .example_images_metadata import MetadataUpdater, update_cache_from_metadata
|
||||||
from ..utils.metadata_manager import MetadataManager
|
from ..utils.metadata_manager import MetadataManager
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
@@ -113,6 +113,26 @@ class ExampleImagesProcessor:
|
|||||||
message = str(error).lower()
|
message = str(error).lower()
|
||||||
return '404' in message or 'file not found' in message
|
return '404' in message or 'file not found' in message
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _example_image_file_exists(model_dir: str, index: int, media_type_hint: str | None = None) -> bool:
|
||||||
|
"""Return True when the file that would be written for a media index already exists.
|
||||||
|
|
||||||
|
The final filename (``image_{index}{extension}``) depends on the downloaded
|
||||||
|
content, so the extension cannot be known ahead of time. The post-download
|
||||||
|
check skips the write when the exact target file exists; this pre-check
|
||||||
|
approximates that with the candidate extensions for the media type (videos
|
||||||
|
only when the metadata hints at a video) so the network request is avoided
|
||||||
|
for files that already exist on disk.
|
||||||
|
"""
|
||||||
|
if media_type_hint == "video":
|
||||||
|
extensions = SUPPORTED_MEDIA_EXTENSIONS['videos']
|
||||||
|
else:
|
||||||
|
extensions = SUPPORTED_MEDIA_EXTENSIONS['images']
|
||||||
|
return any(
|
||||||
|
os.path.exists(os.path.join(model_dir, f"image_{index}{ext}"))
|
||||||
|
for ext in extensions
|
||||||
|
)
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
async def download_model_images(model_hash, model_name, model_images, model_dir, optimize, downloader):
|
async def download_model_images(model_hash, model_name, model_images, model_dir, optimize, downloader):
|
||||||
"""Download images for a single model
|
"""Download images for a single model
|
||||||
@@ -139,7 +159,12 @@ class ExampleImagesProcessor:
|
|||||||
original_url = image_url
|
original_url = image_url
|
||||||
if optimize and 'civitai.com' in image_url:
|
if optimize and 'civitai.com' in image_url:
|
||||||
image_url = ExampleImagesProcessor.get_civitai_optimized_url(image_url)
|
image_url = ExampleImagesProcessor.get_civitai_optimized_url(image_url)
|
||||||
|
|
||||||
|
# Skip the download when the file already exists on disk
|
||||||
|
if ExampleImagesProcessor._example_image_file_exists(model_dir, i, image.get("type")):
|
||||||
|
logger.debug("File already exists, skipping download for %s", image_url)
|
||||||
|
continue
|
||||||
|
|
||||||
# Download the file first to determine the actual file type
|
# Download the file first to determine the actual file type
|
||||||
try:
|
try:
|
||||||
logger.debug(f"Downloading media file {i} for {model_name}")
|
logger.debug(f"Downloading media file {i} for {model_name}")
|
||||||
@@ -229,6 +254,11 @@ class ExampleImagesProcessor:
|
|||||||
if optimize and 'civitai.com' in image_url:
|
if optimize and 'civitai.com' in image_url:
|
||||||
image_url = ExampleImagesProcessor.get_civitai_optimized_url(image_url)
|
image_url = ExampleImagesProcessor.get_civitai_optimized_url(image_url)
|
||||||
|
|
||||||
|
# Skip the download when the file already exists on disk
|
||||||
|
if ExampleImagesProcessor._example_image_file_exists(model_dir, i, image.get("type")):
|
||||||
|
logger.debug("File already exists, skipping download for %s", image_url)
|
||||||
|
continue
|
||||||
|
|
||||||
async def _attempt_download() -> tuple:
|
async def _attempt_download() -> tuple:
|
||||||
logger.debug("Downloading media file %s for %s", i, model_name)
|
logger.debug("Downloading media file %s for %s", i, model_name)
|
||||||
return await downloader.download_to_memory(
|
return await downloader.download_to_memory(
|
||||||
@@ -644,7 +674,7 @@ class ExampleImagesProcessor:
|
|||||||
}, status=500)
|
}, status=500)
|
||||||
|
|
||||||
# Update cache
|
# Update cache
|
||||||
await scanner.update_single_model_cache(file_path, file_path, model_data)
|
await update_cache_from_metadata(scanner, file_path, model_data)
|
||||||
|
|
||||||
# Get regular images array (might be None)
|
# Get regular images array (might be None)
|
||||||
regular_images = civitai_data.get('images', [])
|
regular_images = civitai_data.get('images', [])
|
||||||
@@ -759,7 +789,7 @@ class ExampleImagesProcessor:
|
|||||||
model_copy = model_data.copy()
|
model_copy = model_data.copy()
|
||||||
model_copy.pop('folder', None)
|
model_copy.pop('folder', None)
|
||||||
await MetadataManager.save_metadata(file_path, model_copy)
|
await MetadataManager.save_metadata(file_path, model_copy)
|
||||||
await scanner.update_single_model_cache(file_path, file_path, model_data)
|
await update_cache_from_metadata(scanner, file_path, model_copy)
|
||||||
|
|
||||||
return web.json_response({
|
return web.json_response({
|
||||||
'success': True,
|
'success': True,
|
||||||
|
|||||||
+37
-1
@@ -1,9 +1,10 @@
|
|||||||
|
import functools
|
||||||
import json
|
import json
|
||||||
import logging
|
import logging
|
||||||
import os
|
import os
|
||||||
import struct
|
import struct
|
||||||
from io import BytesIO
|
from io import BytesIO
|
||||||
from typing import Any, Optional
|
from typing import Any, Optional, Tuple
|
||||||
|
|
||||||
import piexif
|
import piexif
|
||||||
from PIL import Image, PngImagePlugin
|
from PIL import Image, PngImagePlugin
|
||||||
@@ -17,6 +18,22 @@ except ImportError:
|
|||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
|
@functools.lru_cache(maxsize=2048)
|
||||||
|
def _get_image_dimensions_cached(path: str, _mtime_ns: int, _size: int) -> Optional[Tuple[int, int]]:
|
||||||
|
"""Return ``(width, height)`` for ``path``, or ``None`` on any failure.
|
||||||
|
|
||||||
|
The ``_mtime_ns`` and ``_size`` arguments are part of the cache key only;
|
||||||
|
they invalidate the entry when the file is replaced with a new image, so a
|
||||||
|
stale preview never serves outdated dimensions.
|
||||||
|
"""
|
||||||
|
try:
|
||||||
|
with Image.open(path) as img:
|
||||||
|
return img.size
|
||||||
|
except Exception:
|
||||||
|
return None
|
||||||
|
|
||||||
|
|
||||||
class ExifUtils:
|
class ExifUtils:
|
||||||
"""Utility functions for working with EXIF data in images"""
|
"""Utility functions for working with EXIF data in images"""
|
||||||
|
|
||||||
@@ -422,6 +439,25 @@ class ExifUtils:
|
|||||||
# Metadata is in the middle of the string
|
# Metadata is in the middle of the string
|
||||||
return user_comment[:recipe_marker_index] + user_comment[next_line_index:]
|
return user_comment[:recipe_marker_index] + user_comment[next_line_index:]
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def get_image_dimensions(image_path: str) -> Optional[Tuple[int, int]]:
|
||||||
|
"""Return ``(width, height)`` for an image, or ``None`` if unavailable.
|
||||||
|
|
||||||
|
Video containers (``.mp4``/``.webm``/``.avi``) and formats PIL cannot
|
||||||
|
read (``.avif``/``.jxl``) return ``None`` before PIL is invoked.
|
||||||
|
Missing or corrupt files return ``None``. Never raises.
|
||||||
|
"""
|
||||||
|
try:
|
||||||
|
ext = os.path.splitext(image_path)[1].lower()
|
||||||
|
if ext in ('.mp4', '.webm', '.avi', '.avif', '.jxl'):
|
||||||
|
return None
|
||||||
|
stat = os.stat(image_path)
|
||||||
|
return _get_image_dimensions_cached(
|
||||||
|
image_path, stat.st_mtime_ns, stat.st_size
|
||||||
|
)
|
||||||
|
except Exception:
|
||||||
|
return None
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def optimize_image(image_data, target_width=250, format='webp', quality=85, preserve_metadata=False):
|
def optimize_image(image_data, target_width=250, format='webp', quality=85, preserve_metadata=False):
|
||||||
"""
|
"""
|
||||||
|
|||||||
+48
-1
@@ -1,7 +1,7 @@
|
|||||||
from difflib import SequenceMatcher
|
from difflib import SequenceMatcher
|
||||||
import os
|
import os
|
||||||
import re
|
import re
|
||||||
from typing import Dict
|
from typing import Any, Dict, List, Optional
|
||||||
from ..services.service_registry import ServiceRegistry
|
from ..services.service_registry import ServiceRegistry
|
||||||
from ..config import config
|
from ..config import config
|
||||||
from ..services.settings_manager import get_settings_manager
|
from ..services.settings_manager import get_settings_manager
|
||||||
@@ -294,6 +294,53 @@ def _format_model_name_for_comfyui(file_path: str, model_roots: list) -> str:
|
|||||||
return os.path.basename(file_path)
|
return os.path.basename(file_path)
|
||||||
|
|
||||||
|
|
||||||
|
def model_patcher_to_name(model_patcher: Any) -> Optional[str]:
|
||||||
|
"""Extract a ComfyUI-style model name from a MODEL (ModelPatcher) object.
|
||||||
|
|
||||||
|
Core ComfyUI loaders record the absolute weight file path on the patcher's
|
||||||
|
``cached_patcher_init`` attribute:
|
||||||
|
- load_checkpoint_guess_config -> (fn, (ckpt_path, ...), index)
|
||||||
|
- load_diffusion_model -> (fn, (unet_path, model_options))
|
||||||
|
Patcher clones (LoRA loaders, model merges, ...) preserve the attribute,
|
||||||
|
so the name is recoverable anywhere downstream of a core loader — including
|
||||||
|
from LoRA Manager's own loaders (CheckpointLoaderLM / UNETLoaderLM), which
|
||||||
|
call the same core load functions.
|
||||||
|
|
||||||
|
The absolute path is converted to the ComfyUI-style relative name used by
|
||||||
|
the metadata pipeline (covering standard ComfyUI roots and LoRA Manager
|
||||||
|
extra folder paths).
|
||||||
|
|
||||||
|
Returns None when the path cannot be recovered (e.g. third-party loaders
|
||||||
|
that never set ``cached_patcher_init``).
|
||||||
|
"""
|
||||||
|
init = getattr(model_patcher, "cached_patcher_init", None)
|
||||||
|
if not isinstance(init, (tuple, list)) or len(init) < 2:
|
||||||
|
return None
|
||||||
|
args = init[1]
|
||||||
|
abs_path = args[0] if args else None
|
||||||
|
if not isinstance(abs_path, str) or not abs_path:
|
||||||
|
return None
|
||||||
|
return _abs_model_path_to_name(abs_path)
|
||||||
|
|
||||||
|
|
||||||
|
def _abs_model_path_to_name(abs_path: str) -> str:
|
||||||
|
"""Convert an absolute model path to a ComfyUI-style relative name.
|
||||||
|
|
||||||
|
Tries standard ComfyUI model roots plus LoRA Manager extra folder paths;
|
||||||
|
falls back to the bare filename.
|
||||||
|
"""
|
||||||
|
try:
|
||||||
|
roots: List[str] = list(config.base_models_roots or [])
|
||||||
|
roots.extend(config.extra_checkpoints_roots or [])
|
||||||
|
roots.extend(config.extra_unet_roots or [])
|
||||||
|
formatted = _format_model_name_for_comfyui(abs_path, roots)
|
||||||
|
if formatted:
|
||||||
|
return formatted
|
||||||
|
except Exception:
|
||||||
|
pass
|
||||||
|
return os.path.basename(abs_path)
|
||||||
|
|
||||||
|
|
||||||
def fuzzy_match(text: str, pattern: str, threshold: float = 0.85) -> bool:
|
def fuzzy_match(text: str, pattern: str, threshold: float = 0.85) -> bool:
|
||||||
"""
|
"""
|
||||||
Check if text matches pattern using fuzzy matching.
|
Check if text matches pattern using fuzzy matching.
|
||||||
|
|||||||
@@ -31,6 +31,10 @@
|
|||||||
overflow: hidden;
|
overflow: hidden;
|
||||||
}
|
}
|
||||||
|
|
||||||
|
.card-grid.masonry-layout .model-card {
|
||||||
|
aspect-ratio: auto;
|
||||||
|
}
|
||||||
|
|
||||||
.model-card:hover {
|
.model-card:hover {
|
||||||
transform: translateY(-2px);
|
transform: translateY(-2px);
|
||||||
box-shadow: var(--shadow-md);
|
box-shadow: var(--shadow-md);
|
||||||
|
|||||||
@@ -0,0 +1,67 @@
|
|||||||
|
/* Batch Download Summary Modal — component styles only.
|
||||||
|
Stat cards and failure table styles are shared with the metadata refresh
|
||||||
|
result modal (metadata-refresh-result.css) and are not redefined here. */
|
||||||
|
|
||||||
|
.download-batch-summary-modal {
|
||||||
|
max-width: 700px;
|
||||||
|
}
|
||||||
|
|
||||||
|
.summary-header {
|
||||||
|
display: flex;
|
||||||
|
align-items: center;
|
||||||
|
gap: var(--space-2);
|
||||||
|
margin: var(--space-2) 0;
|
||||||
|
}
|
||||||
|
|
||||||
|
.summary-header i {
|
||||||
|
font-size: 1.4em;
|
||||||
|
}
|
||||||
|
|
||||||
|
.summary-header.success i {
|
||||||
|
color: var(--color-success);
|
||||||
|
}
|
||||||
|
|
||||||
|
.summary-header.warning i {
|
||||||
|
color: var(--color-warning);
|
||||||
|
}
|
||||||
|
|
||||||
|
.summary-header.error i {
|
||||||
|
color: var(--color-error);
|
||||||
|
}
|
||||||
|
|
||||||
|
.summary-title {
|
||||||
|
font-weight: var(--weight-semibold);
|
||||||
|
color: var(--lora-text);
|
||||||
|
}
|
||||||
|
|
||||||
|
.summary-hint {
|
||||||
|
margin-left: auto;
|
||||||
|
font-size: var(--text-xs);
|
||||||
|
color: var(--text-secondary);
|
||||||
|
}
|
||||||
|
|
||||||
|
.btn-retry {
|
||||||
|
display: inline-flex;
|
||||||
|
align-items: center;
|
||||||
|
gap: var(--space-1);
|
||||||
|
background: var(--lora-accent, #4f46e5);
|
||||||
|
color: #fff;
|
||||||
|
border: none;
|
||||||
|
border-radius: var(--border-radius-sm);
|
||||||
|
padding: var(--space-2) var(--space-3);
|
||||||
|
cursor: pointer;
|
||||||
|
font-weight: var(--weight-semibold);
|
||||||
|
}
|
||||||
|
|
||||||
|
.btn-retry:hover {
|
||||||
|
background: var(--lora-accent-hover, #4338ca);
|
||||||
|
}
|
||||||
|
|
||||||
|
.failure-link {
|
||||||
|
color: var(--lora-accent, #4f46e5);
|
||||||
|
text-decoration: none;
|
||||||
|
}
|
||||||
|
|
||||||
|
.failure-link:hover {
|
||||||
|
text-decoration: underline;
|
||||||
|
}
|
||||||
@@ -1,4 +1,35 @@
|
|||||||
/* Import Modal Styles */
|
/* Import Modal Styles */
|
||||||
|
|
||||||
|
/* Sticky footer layout: fixed header, scrollable step content, pinned action buttons.
|
||||||
|
Ensures Back/Import buttons stay visible on short viewports (e.g. 1080p or 150% zoom). */
|
||||||
|
#importModal .modal-content {
|
||||||
|
display: flex;
|
||||||
|
flex-direction: column;
|
||||||
|
overflow: hidden; /* The active step scrolls instead of the whole modal */
|
||||||
|
}
|
||||||
|
|
||||||
|
#importModal .modal-header {
|
||||||
|
flex-shrink: 0;
|
||||||
|
}
|
||||||
|
|
||||||
|
#importModal .import-step {
|
||||||
|
flex: 1 1 auto;
|
||||||
|
min-height: 0; /* Allow the step to shrink and scroll within the flex container */
|
||||||
|
overflow-y: auto;
|
||||||
|
overflow-x: hidden;
|
||||||
|
scrollbar-gutter: stable;
|
||||||
|
}
|
||||||
|
|
||||||
|
#importModal .import-step .modal-actions {
|
||||||
|
position: sticky;
|
||||||
|
bottom: 0;
|
||||||
|
z-index: 1;
|
||||||
|
background: var(--lora-surface);
|
||||||
|
border-top: 1px solid var(--lora-border);
|
||||||
|
padding-top: var(--space-2);
|
||||||
|
padding-bottom: var(--space-1);
|
||||||
|
}
|
||||||
|
|
||||||
.import-step {
|
.import-step {
|
||||||
margin: var(--space-2) 0;
|
margin: var(--space-2) 0;
|
||||||
transition: none !important;
|
transition: none !important;
|
||||||
|
|||||||
@@ -145,13 +145,14 @@
|
|||||||
position: fixed;
|
position: fixed;
|
||||||
right: 20px;
|
right: 20px;
|
||||||
top: 50px; /* Position below header */
|
top: 50px; /* Position below header */
|
||||||
width: 366px;
|
width: 420px;
|
||||||
background-color: var(--card-bg);
|
background-color: var(--card-bg);
|
||||||
border: 1px solid var(--border-color);
|
border: 1px solid var(--border-color);
|
||||||
border-radius: var(--border-radius-base);
|
border-radius: var(--border-radius-base);
|
||||||
box-shadow: var(--shadow-md);
|
box-shadow: var(--shadow-md);
|
||||||
z-index: var(--z-overlay);
|
z-index: var(--z-overlay);
|
||||||
padding: 16px;
|
padding: 16px;
|
||||||
|
box-sizing: border-box; /* Include padding in max-height calculation */
|
||||||
transition: transform 0.3s ease, opacity 0.3s ease;
|
transition: transform 0.3s ease, opacity 0.3s ease;
|
||||||
transform-origin: top right;
|
transform-origin: top right;
|
||||||
max-height: calc(100vh - 70px); /* Adjusted for header height */
|
max-height: calc(100vh - 70px); /* Adjusted for header height */
|
||||||
@@ -563,7 +564,7 @@
|
|||||||
align-items: center;
|
align-items: center;
|
||||||
gap: 6px;
|
gap: 6px;
|
||||||
white-space: nowrap;
|
white-space: nowrap;
|
||||||
max-width: 120px; /* Prevent long names from breaking layout */
|
max-width: 180px; /* Prevent long names from breaking layout */
|
||||||
overflow: hidden;
|
overflow: hidden;
|
||||||
text-overflow: ellipsis;
|
text-overflow: ellipsis;
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -41,6 +41,7 @@
|
|||||||
@import 'components/sidebar.css'; /* Add sidebar component */
|
@import 'components/sidebar.css'; /* Add sidebar component */
|
||||||
@import 'components/media-viewer.css';
|
@import 'components/media-viewer.css';
|
||||||
@import 'components/metadata-refresh-result.css';
|
@import 'components/metadata-refresh-result.css';
|
||||||
|
@import 'components/download-batch-summary.css';
|
||||||
|
|
||||||
.initialization-notice {
|
.initialization-notice {
|
||||||
display: flex;
|
display: flex;
|
||||||
|
|||||||
@@ -184,7 +184,8 @@ export const DOWNLOAD_ENDPOINTS = {
|
|||||||
downloadGet: '/api/lm/download-model-get',
|
downloadGet: '/api/lm/download-model-get',
|
||||||
cancelGet: '/api/lm/cancel-download-get',
|
cancelGet: '/api/lm/cancel-download-get',
|
||||||
progress: '/api/lm/download-progress',
|
progress: '/api/lm/download-progress',
|
||||||
exampleImages: '/api/lm/force-download-example-images' // New endpoint for downloading example images
|
exampleImages: '/api/lm/force-download-example-images', // Re-process example images ignoring previous status
|
||||||
|
exampleImagesMissing: '/api/lm/download-example-images' // Download only missing example images
|
||||||
};
|
};
|
||||||
|
|
||||||
// Hugging Face API endpoints
|
// Hugging Face API endpoints
|
||||||
|
|||||||
@@ -1641,7 +1641,7 @@ export class BaseModelApiClient {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
async downloadExampleImages(modelHashes, modelTypes = null) {
|
async downloadExampleImages(modelHashes, modelTypes = null, { force = true } = {}) {
|
||||||
let ws = null;
|
let ws = null;
|
||||||
|
|
||||||
await state.loadingManager.showWithProgress(async (loading) => {
|
await state.loadingManager.showWithProgress(async (loading) => {
|
||||||
@@ -1700,8 +1700,13 @@ export class BaseModelApiClient {
|
|||||||
// Determine optimize setting
|
// Determine optimize setting
|
||||||
const optimize = state.global?.settings?.optimize_example_images ?? true;
|
const optimize = state.global?.settings?.optimize_example_images ?? true;
|
||||||
|
|
||||||
|
// force=false routes to the regular endpoint, which skips already-processed models
|
||||||
|
const endpoint = force
|
||||||
|
? DOWNLOAD_ENDPOINTS.exampleImages
|
||||||
|
: DOWNLOAD_ENDPOINTS.exampleImagesMissing;
|
||||||
|
|
||||||
// Make the API request to start the download process
|
// Make the API request to start the download process
|
||||||
const response = await fetch(DOWNLOAD_ENDPOINTS.exampleImages, {
|
const response = await fetch(endpoint, {
|
||||||
method: 'POST',
|
method: 'POST',
|
||||||
headers: {
|
headers: {
|
||||||
'Content-Type': 'application/json'
|
'Content-Type': 'application/json'
|
||||||
@@ -1710,6 +1715,7 @@ export class BaseModelApiClient {
|
|||||||
model_hashes: modelHashes,
|
model_hashes: modelHashes,
|
||||||
output_dir: outputDir,
|
output_dir: outputDir,
|
||||||
optimize: optimize,
|
optimize: optimize,
|
||||||
|
force: force,
|
||||||
model_types: modelTypes || [this.apiConfig.config.singularName]
|
model_types: modelTypes || [this.apiConfig.config.singularName]
|
||||||
})
|
})
|
||||||
});
|
});
|
||||||
|
|||||||
@@ -137,11 +137,10 @@ export class BulkContextMenu extends BaseContextMenu {
|
|||||||
downloadMissingLorasItem.style.display = currentModelType === 'recipes' ? 'flex' : 'none';
|
downloadMissingLorasItem.style.display = currentModelType === 'recipes' ? 'flex' : 'none';
|
||||||
}
|
}
|
||||||
|
|
||||||
const downloadExampleImagesItem = this.menu.querySelector('[data-action="download-example-images"]');
|
const downloadExampleImagesSubmenu = this.menu.querySelector('[data-has-submenu="download-example-images"]');
|
||||||
if (downloadExampleImagesItem) {
|
if (downloadExampleImagesSubmenu) {
|
||||||
// Show on model pages (loras, checkpoints, embeddings), hide on recipes
|
// Show on model pages (loras, checkpoints, embeddings), hide on recipes
|
||||||
const modelPages = ['loras', 'checkpoints', 'embeddings'];
|
downloadExampleImagesSubmenu.style.display = ['loras', 'checkpoints', 'embeddings'].includes(currentModelType) ? 'flex' : 'none';
|
||||||
downloadExampleImagesItem.style.display = modelPages.includes(currentModelType) ? 'flex' : 'none';
|
|
||||||
}
|
}
|
||||||
|
|
||||||
const skipMetadataRefreshItem = this.menu.querySelector('[data-action="skip-metadata-refresh"]');
|
const skipMetadataRefreshItem = this.menu.querySelector('[data-action="skip-metadata-refresh"]');
|
||||||
@@ -294,8 +293,11 @@ export class BulkContextMenu extends BaseContextMenu {
|
|||||||
case 'download-missing-loras':
|
case 'download-missing-loras':
|
||||||
this.handleDownloadMissingLoras();
|
this.handleDownloadMissingLoras();
|
||||||
break;
|
break;
|
||||||
|
case 'download-missing-example-images':
|
||||||
|
this.handleDownloadExampleImages({ force: false });
|
||||||
|
break;
|
||||||
case 'download-example-images':
|
case 'download-example-images':
|
||||||
this.handleDownloadExampleImages();
|
this.handleDownloadExampleImages({ force: true });
|
||||||
break;
|
break;
|
||||||
case 'clear':
|
case 'clear':
|
||||||
bulkManager.clearSelection();
|
bulkManager.clearSelection();
|
||||||
@@ -340,7 +342,7 @@ export class BulkContextMenu extends BaseContextMenu {
|
|||||||
await bulkMissingLoraDownloadManager.downloadMissingLoras(selectedRecipes);
|
await bulkMissingLoraDownloadManager.downloadMissingLoras(selectedRecipes);
|
||||||
}
|
}
|
||||||
|
|
||||||
async handleDownloadExampleImages() {
|
async handleDownloadExampleImages({ force = true } = {}) {
|
||||||
if (state.selectedModels.size === 0) {
|
if (state.selectedModels.size === 0) {
|
||||||
return;
|
return;
|
||||||
}
|
}
|
||||||
@@ -361,7 +363,7 @@ export class BulkContextMenu extends BaseContextMenu {
|
|||||||
|
|
||||||
try {
|
try {
|
||||||
const apiClient = getModelApiClient();
|
const apiClient = getModelApiClient();
|
||||||
await apiClient.downloadExampleImages([...hashes]);
|
await apiClient.downloadExampleImages([...hashes], null, { force });
|
||||||
} catch (error) {
|
} catch (error) {
|
||||||
console.error('Bulk download example images failed:', error);
|
console.error('Bulk download example images failed:', error);
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -347,7 +347,10 @@ export const ModelContextMenuMixin = {
|
|||||||
openExampleImagesFolder(this.currentCard.dataset.sha256);
|
openExampleImagesFolder(this.currentCard.dataset.sha256);
|
||||||
return true;
|
return true;
|
||||||
case 'download-examples':
|
case 'download-examples':
|
||||||
this.downloadExampleImages();
|
this.downloadExampleImages(false);
|
||||||
|
return true;
|
||||||
|
case 'download-examples-force':
|
||||||
|
this.downloadExampleImages(true);
|
||||||
return true;
|
return true;
|
||||||
case 'civitai':
|
case 'civitai':
|
||||||
if (this.currentCard.dataset.from_civitai === 'true') {
|
if (this.currentCard.dataset.from_civitai === 'true') {
|
||||||
@@ -378,7 +381,7 @@ export const ModelContextMenuMixin = {
|
|||||||
},
|
},
|
||||||
|
|
||||||
// Download example images method
|
// Download example images method
|
||||||
async downloadExampleImages() {
|
async downloadExampleImages(force = false) {
|
||||||
const modelHash = this.currentCard.dataset.sha256;
|
const modelHash = this.currentCard.dataset.sha256;
|
||||||
if (!modelHash) {
|
if (!modelHash) {
|
||||||
showToast('toast.contextMenu.missingHash', {}, 'error');
|
showToast('toast.contextMenu.missingHash', {}, 'error');
|
||||||
@@ -387,7 +390,7 @@ export const ModelContextMenuMixin = {
|
|||||||
|
|
||||||
try {
|
try {
|
||||||
const apiClient = getModelApiClient();
|
const apiClient = getModelApiClient();
|
||||||
await apiClient.downloadExampleImages([modelHash]);
|
await apiClient.downloadExampleImages([modelHash], null, { force });
|
||||||
} catch (error) {
|
} catch (error) {
|
||||||
console.error('Error downloading example images:', error);
|
console.error('Error downloading example images:', error);
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -0,0 +1,340 @@
|
|||||||
|
import { translate } from '../utils/i18nHelpers.js';
|
||||||
|
import { showToast, openHuggingFace } from '../utils/uiHelpers.js';
|
||||||
|
|
||||||
|
/**
|
||||||
|
* Escape HTML entities in a string to prevent injection when interpolating into innerHTML.
|
||||||
|
* Safe for both text content and attribute values (quotes are escaped too).
|
||||||
|
* @param {string} str - The string to escape
|
||||||
|
* @returns {string} - The escaped string
|
||||||
|
*/
|
||||||
|
function _escapeHtml(str) {
|
||||||
|
if (!str) return '';
|
||||||
|
const div = document.createElement('div');
|
||||||
|
div.textContent = str;
|
||||||
|
return div.innerHTML.replace(/"/g, '"').replace(/'/g, ''');
|
||||||
|
}
|
||||||
|
|
||||||
|
/**
|
||||||
|
* Resolve the display name of a failed download entry.
|
||||||
|
* Prefers the resolved name carried on the entry, then known item fields,
|
||||||
|
* then derives a name from the item URL as a last resort.
|
||||||
|
* @param {Object} entry - The failed entry ({ item, error, name? })
|
||||||
|
* @returns {string} - The best available display name
|
||||||
|
*/
|
||||||
|
function _resolveItemName(entry) {
|
||||||
|
if (entry?.name) {
|
||||||
|
return entry.name;
|
||||||
|
}
|
||||||
|
const item = entry?.item ?? entry;
|
||||||
|
const direct = item?.displayName || item?.name || item?.file_name || item?.filename || item?.selectedVersion?.name;
|
||||||
|
if (direct) {
|
||||||
|
return direct;
|
||||||
|
}
|
||||||
|
if (item?.url) {
|
||||||
|
try {
|
||||||
|
const segments = new URL(item.url).pathname.split('/').filter(Boolean);
|
||||||
|
if (segments.length > 0) {
|
||||||
|
return decodeURIComponent(segments[segments.length - 1]);
|
||||||
|
}
|
||||||
|
} catch (e) {
|
||||||
|
// Unparseable URL — fall through to 'Unknown'
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return 'Unknown';
|
||||||
|
}
|
||||||
|
|
||||||
|
/**
|
||||||
|
* Resolve the URL to open for a failed item — always the original item URL.
|
||||||
|
* @param {Object} item - The failed item payload
|
||||||
|
* @returns {string|null} - A URL string, or null when nothing is available
|
||||||
|
*/
|
||||||
|
function _resolveItemUrl(item) {
|
||||||
|
return item?.url || null;
|
||||||
|
}
|
||||||
|
|
||||||
|
/**
|
||||||
|
* Format a raw failure error into a concise human-readable message.
|
||||||
|
* Unwraps JSON envelopes and extracts HTTP status/body details when present.
|
||||||
|
* @param {*} error - The raw error (usually a string)
|
||||||
|
* @returns {string} - The formatted error message
|
||||||
|
*/
|
||||||
|
function _formatError(error) {
|
||||||
|
if (!error) {
|
||||||
|
return 'Unknown error';
|
||||||
|
}
|
||||||
|
let base = typeof error === 'string' ? error : String(error);
|
||||||
|
|
||||||
|
// Unwrap JSON envelope: { "success": false, "error": "...", ... }
|
||||||
|
try {
|
||||||
|
const parsed = JSON.parse(base);
|
||||||
|
if (parsed && typeof parsed.error === 'string' && parsed.error) {
|
||||||
|
base = parsed.error;
|
||||||
|
}
|
||||||
|
} catch (e) {
|
||||||
|
// Not a JSON envelope — keep the raw string
|
||||||
|
}
|
||||||
|
|
||||||
|
// Extract HTTP status and JSON body details, e.g. "status=403 body={...}"
|
||||||
|
let result = base;
|
||||||
|
const statusMatch = base.match(/status=(\d{3})/);
|
||||||
|
const bodyMatch = base.match(/body=(\{.*\})/s);
|
||||||
|
if (bodyMatch) {
|
||||||
|
try {
|
||||||
|
const body = JSON.parse(bodyMatch[1]);
|
||||||
|
const detail = (typeof body?.message === 'string' && body.message)
|
||||||
|
|| (typeof body?.error === 'string' && body.error)
|
||||||
|
|| null;
|
||||||
|
if (detail) {
|
||||||
|
const status = statusMatch ? statusMatch[1] : null;
|
||||||
|
result = `${status ? `HTTP ${status} — ` : ''}${detail}`;
|
||||||
|
}
|
||||||
|
} catch (e) {
|
||||||
|
// Body is not valid JSON — keep the base string
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Truncate overly long messages
|
||||||
|
if (result.length > 220) {
|
||||||
|
result = result.slice(0, 220) + '…';
|
||||||
|
}
|
||||||
|
return result;
|
||||||
|
}
|
||||||
|
|
||||||
|
/**
|
||||||
|
* Build a plain-text report of the batch download results.
|
||||||
|
* @param {number} total - Total number of models attempted
|
||||||
|
* @param {number} completed - Number of models successfully downloaded
|
||||||
|
* @param {Array} failedItems - Array of failed items ({ item, error })
|
||||||
|
* @returns {string} - The report text
|
||||||
|
*/
|
||||||
|
function _buildReportText(total, completed, failedItems) {
|
||||||
|
const lines = [
|
||||||
|
'=== Batch Download Report ===',
|
||||||
|
`Date: ${new Date().toLocaleString()}`,
|
||||||
|
`Total: ${total}`,
|
||||||
|
`Successfully downloaded: ${completed}`,
|
||||||
|
`Failed: ${failedItems.length}`,
|
||||||
|
'',
|
||||||
|
];
|
||||||
|
if (failedItems.length > 0) {
|
||||||
|
lines.push('--- Failed Items ---');
|
||||||
|
failedItems.forEach((entry, i) => {
|
||||||
|
const name = _resolveItemName(entry);
|
||||||
|
const error = _formatError(entry?.error);
|
||||||
|
lines.push(`${i + 1}. ${name} — ${error}`);
|
||||||
|
const itemUrl = _resolveItemUrl(entry?.item ?? entry);
|
||||||
|
if (itemUrl) {
|
||||||
|
lines.push(` URL: ${itemUrl}`);
|
||||||
|
}
|
||||||
|
});
|
||||||
|
lines.push('');
|
||||||
|
}
|
||||||
|
lines.push('====================');
|
||||||
|
return lines.join('\n');
|
||||||
|
}
|
||||||
|
|
||||||
|
/**
|
||||||
|
* Handle a successful clipboard write: confirm via toast and briefly swap the
|
||||||
|
* trigger button to a "Copied!" state.
|
||||||
|
* @param {HTMLElement|null} btn - The button that triggered the copy action
|
||||||
|
*/
|
||||||
|
function _onCopyReportSuccess(btn) {
|
||||||
|
showToast('toast.api.copiedToClipboard', {}, 'success');
|
||||||
|
if (btn) {
|
||||||
|
const origHTML = btn.innerHTML;
|
||||||
|
btn.innerHTML = '<i class="fas fa-check"></i> Copied!';
|
||||||
|
setTimeout(() => { btn.innerHTML = origHTML; }, 2000);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
/**
|
||||||
|
* Fallback for environments without the async Clipboard API (e.g. insecure
|
||||||
|
* contexts over LAN http where `navigator.clipboard` is undefined): copy via a
|
||||||
|
* hidden textarea and `document.execCommand('copy')`.
|
||||||
|
* @param {string} text - The report text to copy
|
||||||
|
*/
|
||||||
|
function _copyReportWithExecCommand(text) {
|
||||||
|
const textarea = document.createElement('textarea');
|
||||||
|
textarea.value = text;
|
||||||
|
document.body.appendChild(textarea);
|
||||||
|
textarea.select();
|
||||||
|
document.execCommand('copy');
|
||||||
|
document.body.removeChild(textarea);
|
||||||
|
showToast('toast.api.copiedToClipboard', {}, 'success');
|
||||||
|
}
|
||||||
|
|
||||||
|
/**
|
||||||
|
* Copy the batch download report to the clipboard.
|
||||||
|
* Uses the async Clipboard API when available, otherwise falls back to a hidden
|
||||||
|
* textarea + execCommand so the action still works in insecure contexts.
|
||||||
|
* @param {HTMLElement} btn - The button that triggered the copy action
|
||||||
|
* @param {number} total - Total number of models attempted
|
||||||
|
* @param {number} completed - Number of models successfully downloaded
|
||||||
|
* @param {Array} failedItems - Array of failed items
|
||||||
|
*/
|
||||||
|
function _copyReport(btn, total, completed, failedItems) {
|
||||||
|
const text = _buildReportText(total, completed, failedItems);
|
||||||
|
if (navigator.clipboard && typeof navigator.clipboard.writeText === 'function') {
|
||||||
|
navigator.clipboard.writeText(text)
|
||||||
|
.then(() => _onCopyReportSuccess(btn))
|
||||||
|
.catch(() => _copyReportWithExecCommand(text));
|
||||||
|
} else {
|
||||||
|
_copyReportWithExecCommand(text);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
/**
|
||||||
|
* Show the batch download summary modal after a batch download completes.
|
||||||
|
* Mirrors the Metadata Fetch Summary modal lifecycle: the modal element is
|
||||||
|
* appended directly to document.body and removed on close; it is not
|
||||||
|
* registered with ModalManager.
|
||||||
|
* @param {Object} options - Summary options
|
||||||
|
* @param {number} options.total - Total number of models attempted
|
||||||
|
* @param {number} options.completed - Number of models successfully downloaded
|
||||||
|
* @param {Array} options.failedItems - Array of failed items ({ item, error })
|
||||||
|
* @param {Function} options.onRetry - Callback invoked with failedItems to retry the failed subset
|
||||||
|
*/
|
||||||
|
export function showDownloadBatchSummary({ total, completed, failedItems, onRetry }) {
|
||||||
|
const failures = failedItems || [];
|
||||||
|
const failedCount = failures.length;
|
||||||
|
|
||||||
|
// 3-state summary header semantics (mirrors BatchImportManager results header)
|
||||||
|
let headerState;
|
||||||
|
let headerIcon;
|
||||||
|
let headerText;
|
||||||
|
if (completed === 0) {
|
||||||
|
headerState = 'error';
|
||||||
|
headerIcon = 'fa-times-circle';
|
||||||
|
headerText = translate('modals.downloadBatchSummary.failed', {}, 'Download failed');
|
||||||
|
} else if (failedCount > 0) {
|
||||||
|
headerState = 'warning';
|
||||||
|
headerIcon = 'fa-exclamation-circle';
|
||||||
|
headerText = translate('modals.downloadBatchSummary.completedWithErrors', {}, 'Completed with errors');
|
||||||
|
} else {
|
||||||
|
headerState = 'success';
|
||||||
|
headerIcon = 'fa-check-circle';
|
||||||
|
headerText = translate('modals.downloadBatchSummary.successMessage', { count: completed }, 'All ' + completed + ' models downloaded successfully');
|
||||||
|
}
|
||||||
|
|
||||||
|
// Build failure table rows
|
||||||
|
const failureRows = failures.map((entry, i) => {
|
||||||
|
const item = entry?.item ?? entry;
|
||||||
|
const name = _resolveItemName(entry);
|
||||||
|
const itemUrl = _resolveItemUrl(item);
|
||||||
|
const rawError = entry?.error ? String(entry.error) : '';
|
||||||
|
const error = _formatError(entry?.error);
|
||||||
|
const nameCell = itemUrl
|
||||||
|
? `<td class="failure-name"><a href="#" class="failure-link" data-action="open-model" data-index="${i}" title="${_escapeHtml(itemUrl)}">${_escapeHtml(name)}</a></td>`
|
||||||
|
: `<td class="failure-name" title="${_escapeHtml(name)}">${_escapeHtml(name)}</td>`;
|
||||||
|
return `<tr>
|
||||||
|
<td class="failure-index">${i + 1}</td>
|
||||||
|
${nameCell}
|
||||||
|
<td class="failure-error" title="${_escapeHtml(rawError)}">${_escapeHtml(error)}</td>
|
||||||
|
</tr>`;
|
||||||
|
}).join('');
|
||||||
|
|
||||||
|
const modalHtml = `
|
||||||
|
<div id="downloadBatchSummaryModal" class="modal" style="display: block;">
|
||||||
|
<div class="modal-content download-batch-summary-modal">
|
||||||
|
<button class="close" data-action="close-modal">×</button>
|
||||||
|
|
||||||
|
<h2>${translate('modals.downloadBatchSummary.title', {}, 'Batch Download Summary')}</h2>
|
||||||
|
|
||||||
|
<div class="summary-header ${headerState}">
|
||||||
|
<i class="fas ${headerIcon}"></i>
|
||||||
|
<span class="summary-title">${headerText}</span>
|
||||||
|
<span class="summary-hint">${completed}/${total}</span>
|
||||||
|
</div>
|
||||||
|
|
||||||
|
<div class="refresh-summary-stats">
|
||||||
|
<div class="stat-card stat-card-success">
|
||||||
|
<div class="stat-card-body">
|
||||||
|
<span class="stat-card-label">${translate('modals.downloadBatchSummary.statSuccess', {}, 'Success')}</span>
|
||||||
|
<span class="stat-card-value">${completed}</span>
|
||||||
|
</div>
|
||||||
|
</div>
|
||||||
|
<div class="stat-card stat-card-failure">
|
||||||
|
<div class="stat-card-body">
|
||||||
|
<span class="stat-card-label">${translate('modals.downloadBatchSummary.statFailed', {}, 'Failed')}</span>
|
||||||
|
<span class="stat-card-value">${failedCount}</span>
|
||||||
|
</div>
|
||||||
|
</div>
|
||||||
|
<div class="stat-card stat-card-total">
|
||||||
|
<div class="stat-card-body">
|
||||||
|
<span class="stat-card-label">${translate('modals.downloadBatchSummary.statTotal', {}, 'Total')}</span>
|
||||||
|
<span class="stat-card-value">${total}</span>
|
||||||
|
</div>
|
||||||
|
</div>
|
||||||
|
</div>
|
||||||
|
|
||||||
|
${failedCount > 0 ? `
|
||||||
|
<div class="refresh-failures-section">
|
||||||
|
<h4><i class="fas fa-exclamation-triangle"></i> ${translate('modals.downloadBatchSummary.failedItems', { count: failedCount }, 'Failed Items (' + failedCount + ')')}</h4>
|
||||||
|
<div class="failure-table-wrapper">
|
||||||
|
<table class="failure-table">
|
||||||
|
<thead>
|
||||||
|
<tr>
|
||||||
|
<th>#</th>
|
||||||
|
<th>${translate('modals.downloadBatchSummary.columnName', {}, 'Model Name')}</th>
|
||||||
|
<th>${translate('modals.downloadBatchSummary.columnError', {}, 'Error')}</th>
|
||||||
|
</tr>
|
||||||
|
</thead>
|
||||||
|
<tbody>${failureRows}</tbody>
|
||||||
|
</table>
|
||||||
|
</div>
|
||||||
|
</div>
|
||||||
|
` : `
|
||||||
|
<div class="refresh-success-message">
|
||||||
|
<i class="fas fa-check-circle"></i> ${translate('modals.downloadBatchSummary.successMessage', { count: completed }, 'All ' + completed + ' models downloaded successfully')}
|
||||||
|
</div>
|
||||||
|
`}
|
||||||
|
|
||||||
|
<div class="modal-actions">
|
||||||
|
${failedCount > 0 ? `
|
||||||
|
<button class="btn-retry" data-action="retry-failed"><i class="fas fa-redo"></i> ${translate('modals.downloadBatchSummary.retryFailed', { count: failedCount }, 'Retry Failed (' + failedCount + ')')}</button>
|
||||||
|
<button class="secondary-btn" data-action="copy-report"><i class="fas fa-copy"></i> ${translate('modals.downloadBatchSummary.copyReport', {}, 'Copy Report')}</button>
|
||||||
|
` : ''}
|
||||||
|
<button class="cancel-btn" data-action="close-modal">${translate('modals.downloadBatchSummary.close', {}, 'Close')}</button>
|
||||||
|
</div>
|
||||||
|
</div>
|
||||||
|
</div>
|
||||||
|
`;
|
||||||
|
|
||||||
|
const existing = document.getElementById('downloadBatchSummaryModal');
|
||||||
|
if (existing) existing.remove();
|
||||||
|
|
||||||
|
const container = document.createElement('div');
|
||||||
|
container.innerHTML = modalHtml;
|
||||||
|
const modal = container.firstElementChild;
|
||||||
|
document.body.appendChild(modal);
|
||||||
|
|
||||||
|
modal.addEventListener('click', (e) => {
|
||||||
|
const actionEl = e.target.closest('[data-action]');
|
||||||
|
const action = actionEl?.dataset.action;
|
||||||
|
if (!action) return;
|
||||||
|
e.preventDefault();
|
||||||
|
|
||||||
|
switch (action) {
|
||||||
|
case 'close-modal':
|
||||||
|
modal.remove();
|
||||||
|
break;
|
||||||
|
case 'retry-failed':
|
||||||
|
modal.remove();
|
||||||
|
if (typeof onRetry === 'function') {
|
||||||
|
onRetry(failures);
|
||||||
|
}
|
||||||
|
break;
|
||||||
|
case 'copy-report':
|
||||||
|
_copyReport(actionEl, total, completed, failures);
|
||||||
|
break;
|
||||||
|
case 'open-model': {
|
||||||
|
// Keep the modal open; just open the item's original URL in a new tab
|
||||||
|
const entry = failures[Number(actionEl.dataset.index)];
|
||||||
|
const item = entry?.item;
|
||||||
|
if (!item?.url) break;
|
||||||
|
openHuggingFace(item.url);
|
||||||
|
break;
|
||||||
|
}
|
||||||
|
}
|
||||||
|
});
|
||||||
|
}
|
||||||
@@ -2,6 +2,7 @@
|
|||||||
import { showToast } from '../utils/uiHelpers.js';
|
import { showToast } from '../utils/uiHelpers.js';
|
||||||
import { RecipeCard } from './RecipeCard.js';
|
import { RecipeCard } from './RecipeCard.js';
|
||||||
import { state, getCurrentPageState } from '../state/index.js';
|
import { state, getCurrentPageState } from '../state/index.js';
|
||||||
|
import { recreateVirtualScroll } from '../utils/infiniteScroll.js';
|
||||||
|
|
||||||
export class DuplicatesManager {
|
export class DuplicatesManager {
|
||||||
constructor(recipeManager) {
|
constructor(recipeManager) {
|
||||||
@@ -71,7 +72,7 @@ export class DuplicatesManager {
|
|||||||
this.updateSelectedCount();
|
this.updateSelectedCount();
|
||||||
}
|
}
|
||||||
|
|
||||||
exitDuplicateMode() {
|
async exitDuplicateMode() {
|
||||||
this.inDuplicateMode = false;
|
this.inDuplicateMode = false;
|
||||||
this.selectedForDeletion.clear();
|
this.selectedForDeletion.clear();
|
||||||
|
|
||||||
@@ -94,8 +95,16 @@ export class DuplicatesManager {
|
|||||||
recipeGrid.innerHTML = '';
|
recipeGrid.innerHTML = '';
|
||||||
}
|
}
|
||||||
|
|
||||||
// Re-enable virtual scrolling
|
// Re-enable virtual scrolling, or apply a layout switch deferred
|
||||||
state.virtualScroller.enable();
|
// while duplicates mode was active (enabling the old scroller first
|
||||||
|
// would let its pending rAF render repopulate the grid after the
|
||||||
|
// new scroller is created, leaving orphaned/overlapping cards).
|
||||||
|
if (state.pendingLayoutRecreate) {
|
||||||
|
state.pendingLayoutRecreate = false;
|
||||||
|
await recreateVirtualScroll('recipes');
|
||||||
|
} else if (state.virtualScroller) {
|
||||||
|
state.virtualScroller.enable();
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
renderDuplicateGroups() {
|
renderDuplicateGroups() {
|
||||||
|
|||||||
@@ -108,10 +108,20 @@ export class PageControls {
|
|||||||
const sortSelect = document.getElementById('sortSelect');
|
const sortSelect = document.getElementById('sortSelect');
|
||||||
if (sortSelect) {
|
if (sortSelect) {
|
||||||
initSortDropdown(sortSelect);
|
initSortDropdown(sortSelect);
|
||||||
sortSelect.value = this.pageState.sortBy;
|
this.applySortToSelect(this.pageState.sortBy);
|
||||||
sortSelect.addEventListener('change', async (e) => {
|
sortSelect.addEventListener('change', async (e) => {
|
||||||
this.pageState.sortBy = e.target.value;
|
let value = e.target.value;
|
||||||
this.saveSortPreference(e.target.value);
|
if (value.startsWith('random')) {
|
||||||
|
// Every pick of Random reshuffles the list: generate a
|
||||||
|
// fresh seed so the backend keeps a stable order across
|
||||||
|
// paginated requests.
|
||||||
|
value = this._randomizeSortValue();
|
||||||
|
}
|
||||||
|
this.pageState.sortBy = value;
|
||||||
|
this.saveSortPreference(value);
|
||||||
|
// Reset the seeded Random option when switching away from
|
||||||
|
// Random, or re-apply the fresh seed when picking it again.
|
||||||
|
this.applySortToSelect(value);
|
||||||
await this.resetAndReload();
|
await this.resetAndReload();
|
||||||
});
|
});
|
||||||
}
|
}
|
||||||
@@ -312,6 +322,44 @@ export class PageControls {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
/**
|
||||||
|
* Apply a sort value to the native sort <select>, keeping the Random
|
||||||
|
* option's value in sync when the persisted value carries a seed
|
||||||
|
* (e.g. "random:abc123"). Must be used instead of assigning
|
||||||
|
* sortSelect.value directly whenever the value may be a seeded random
|
||||||
|
* sort, otherwise the native select has no matching option.
|
||||||
|
* @param {string} sortValue - Sort value like "name:asc" or "random:<seed>"
|
||||||
|
*/
|
||||||
|
applySortToSelect(sortValue) {
|
||||||
|
const sortSelect = document.getElementById('sortSelect');
|
||||||
|
if (!sortSelect) return;
|
||||||
|
const randomOpt = sortSelect.querySelector('option[value="random"], option[value^="random:"]');
|
||||||
|
if (randomOpt) {
|
||||||
|
randomOpt.value = String(sortValue).startsWith('random') ? sortValue : 'random';
|
||||||
|
}
|
||||||
|
sortSelect.value = sortValue;
|
||||||
|
}
|
||||||
|
|
||||||
|
/**
|
||||||
|
* Generate a fresh seeded random sort value ("random:<seed>") and keep
|
||||||
|
* the native <select> in sync so its value matches the persisted sort
|
||||||
|
* string and the dropdown shows the selected label.
|
||||||
|
* @returns {string} The new sort value, e.g. "random:abc123xyz"
|
||||||
|
*/
|
||||||
|
_randomizeSortValue() {
|
||||||
|
const seed = Math.random().toString(36).slice(2, 12);
|
||||||
|
const value = `random:${seed}`;
|
||||||
|
const sortSelect = document.getElementById('sortSelect');
|
||||||
|
if (sortSelect) {
|
||||||
|
const randomOpt = sortSelect.querySelector('option[value="random"], option[value^="random:"]');
|
||||||
|
if (randomOpt) {
|
||||||
|
randomOpt.value = value;
|
||||||
|
}
|
||||||
|
sortSelect.value = value;
|
||||||
|
}
|
||||||
|
return value;
|
||||||
|
}
|
||||||
|
|
||||||
/**
|
/**
|
||||||
* Load sort preference from storage
|
* Load sort preference from storage
|
||||||
*/
|
*/
|
||||||
@@ -326,10 +374,7 @@ export class PageControls {
|
|||||||
// Handle legacy format conversion
|
// Handle legacy format conversion
|
||||||
const convertedSort = this.convertLegacySortFormat(savedSort);
|
const convertedSort = this.convertLegacySortFormat(savedSort);
|
||||||
this.pageState.sortBy = convertedSort;
|
this.pageState.sortBy = convertedSort;
|
||||||
const sortSelect = document.getElementById('sortSelect');
|
this.applySortToSelect(convertedSort);
|
||||||
if (sortSelect) {
|
|
||||||
sortSelect.value = convertedSort;
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -523,9 +568,9 @@ export class PageControls {
|
|||||||
this.pageState.sortBy = restoredSort;
|
this.pageState.sortBy = restoredSort;
|
||||||
this.saveSortPreference(restoredSort);
|
this.saveSortPreference(restoredSort);
|
||||||
this._removeVlmSortOption();
|
this._removeVlmSortOption();
|
||||||
|
this.applySortToSelect(restoredSort);
|
||||||
const sortSelect = document.getElementById('sortSelect');
|
const sortSelect = document.getElementById('sortSelect');
|
||||||
if (sortSelect) {
|
if (sortSelect) {
|
||||||
sortSelect.value = restoredSort;
|
|
||||||
sortSelect.disabled = false;
|
sortSelect.disabled = false;
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -575,10 +620,7 @@ export class PageControls {
|
|||||||
const savedGroupedSort = getStorageItem(groupedKey);
|
const savedGroupedSort = getStorageItem(groupedKey);
|
||||||
if (savedGroupedSort) {
|
if (savedGroupedSort) {
|
||||||
this.pageState.sortBy = savedGroupedSort;
|
this.pageState.sortBy = savedGroupedSort;
|
||||||
const sortSelect = document.getElementById('sortSelect');
|
this.applySortToSelect(savedGroupedSort);
|
||||||
if (sortSelect) {
|
|
||||||
sortSelect.value = savedGroupedSort;
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
} else {
|
} else {
|
||||||
// Leaving group mode: persist current sort for next time, restore non-group sort
|
// Leaving group mode: persist current sort for next time, restore non-group sort
|
||||||
@@ -586,10 +628,7 @@ export class PageControls {
|
|||||||
const savedNormalSort = getStorageItem(`${this.pageType}_sort`);
|
const savedNormalSort = getStorageItem(`${this.pageType}_sort`);
|
||||||
if (savedNormalSort) {
|
if (savedNormalSort) {
|
||||||
this.pageState.sortBy = savedNormalSort;
|
this.pageState.sortBy = savedNormalSort;
|
||||||
const sortSelect = document.getElementById('sortSelect');
|
this.applySortToSelect(savedNormalSort);
|
||||||
if (sortSelect) {
|
|
||||||
sortSelect.value = savedNormalSort;
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -874,7 +913,7 @@ export class PageControls {
|
|||||||
}
|
}
|
||||||
|
|
||||||
if (sortSelect) {
|
if (sortSelect) {
|
||||||
sortSelect.value = this.pageState.sortBy;
|
this.applySortToSelect(this.pageState.sortBy);
|
||||||
}
|
}
|
||||||
if (searchInput) {
|
if (searchInput) {
|
||||||
searchInput.value = this.pageState.filters?.search || '';
|
searchInput.value = this.pageState.filters?.search || '';
|
||||||
|
|||||||
@@ -96,7 +96,16 @@ export function initSortDropdown(select) {
|
|||||||
};
|
};
|
||||||
|
|
||||||
const choose = (value) => {
|
const choose = (value) => {
|
||||||
if (select.value === value) return;
|
if (select.value === value) {
|
||||||
|
// Re-picking the already-selected option is normally a no-op,
|
||||||
|
// matching native <select> behavior. The seeded Random sort is
|
||||||
|
// the exception: clicking it again should reshuffle, so let the
|
||||||
|
// change handler (PageControls) generate a fresh seed.
|
||||||
|
if (String(value).startsWith('random')) {
|
||||||
|
select.dispatchEvent(new Event('change', { bubbles: true }));
|
||||||
|
}
|
||||||
|
return;
|
||||||
|
}
|
||||||
select.value = value;
|
select.value = value;
|
||||||
select.dispatchEvent(new Event('change', { bubbles: true }));
|
select.dispatchEvent(new Event('change', { bubbles: true }));
|
||||||
};
|
};
|
||||||
@@ -277,9 +286,10 @@ export function initSortDropdown(select) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
// Rebuild the menu when <option>s change (VLM adds/removes a temporary
|
// Rebuild the menu when <option>s change (VLM adds/removes a temporary
|
||||||
// option at runtime).
|
// option at runtime, and the seeded Random sort option gets a new value
|
||||||
|
// attribute each time it is picked).
|
||||||
const observer = new MutationObserver(() => buildMenu());
|
const observer = new MutationObserver(() => buildMenu());
|
||||||
observer.observe(select, { childList: true });
|
observer.observe(select, { childList: true, subtree: true, attributes: true, attributeFilter: ['value'] });
|
||||||
|
|
||||||
buildMenu();
|
buildMenu();
|
||||||
group.dataset.sortReady = '1';
|
group.dataset.sortReady = '1';
|
||||||
|
|||||||
@@ -90,7 +90,7 @@ export class BulkManager {
|
|||||||
moveAll: true,
|
moveAll: true,
|
||||||
autoOrganize: false,
|
autoOrganize: false,
|
||||||
deleteAll: true,
|
deleteAll: true,
|
||||||
setContentRating: false,
|
setContentRating: true,
|
||||||
skipMetadataRefresh: false,
|
skipMetadataRefresh: false,
|
||||||
setFavorite: true,
|
setFavorite: true,
|
||||||
unfavorite: true,
|
unfavorite: true,
|
||||||
@@ -1528,14 +1528,18 @@ export class BulkManager {
|
|||||||
let failureCount = 0;
|
let failureCount = 0;
|
||||||
|
|
||||||
try {
|
try {
|
||||||
const apiClient = getModelApiClient();
|
const isRecipesPage = state.currentPageType === 'recipes';
|
||||||
for (const filePath of targets) {
|
for (const filePath of targets) {
|
||||||
if (cancelled) {
|
if (cancelled) {
|
||||||
showToast('toast.api.operationCancelled', {}, 'info');
|
showToast('toast.api.operationCancelled', {}, 'info');
|
||||||
break;
|
break;
|
||||||
}
|
}
|
||||||
try {
|
try {
|
||||||
await apiClient.saveModelMetadata(filePath, { preview_nsfw_level: level });
|
if (isRecipesPage) {
|
||||||
|
await updateRecipeMetadata(filePath, { preview_nsfw_level: level });
|
||||||
|
} else {
|
||||||
|
await getModelApiClient().saveModelMetadata(filePath, { preview_nsfw_level: level });
|
||||||
|
}
|
||||||
successCount++;
|
successCount++;
|
||||||
} catch (error) {
|
} catch (error) {
|
||||||
failureCount++;
|
failureCount++;
|
||||||
|
|||||||
@@ -6,8 +6,9 @@ import { getModelApiClient, resetAndReload } from '../api/modelApiFactory.js';
|
|||||||
import { getStorageItem, setStorageItem } from '../utils/storageHelpers.js';
|
import { getStorageItem, setStorageItem } from '../utils/storageHelpers.js';
|
||||||
import { FolderTreeManager } from '../components/FolderTreeManager.js';
|
import { FolderTreeManager } from '../components/FolderTreeManager.js';
|
||||||
import { translate } from '../utils/i18nHelpers.js';
|
import { translate } from '../utils/i18nHelpers.js';
|
||||||
import { extractCivitaiModelUrlParts } from '../utils/civitaiUtils.js';
|
import { buildCivitaiUrl, extractCivitaiModelUrlParts, normalizeCivitaiPageHost } from '../utils/civitaiUtils.js';
|
||||||
import { formatFileSize } from '../utils/formatters.js';
|
import { formatFileSize } from '../utils/formatters.js';
|
||||||
|
import { showDownloadBatchSummary } from '../components/DownloadBatchSummaryModal.js';
|
||||||
|
|
||||||
export class DownloadManager {
|
export class DownloadManager {
|
||||||
constructor() {
|
constructor() {
|
||||||
@@ -234,6 +235,9 @@ export class DownloadManager {
|
|||||||
|
|
||||||
if (this.modelVersionId) {
|
if (this.modelVersionId) {
|
||||||
this.currentVersion = this.versions.find(v => v.id.toString() === this.modelVersionId);
|
this.currentVersion = this.versions.find(v => v.id.toString() === this.modelVersionId);
|
||||||
|
} else {
|
||||||
|
// No explicit version id in the URL → default to the latest version (Civitai returns newest first)
|
||||||
|
this.currentVersion = this.versions[0];
|
||||||
}
|
}
|
||||||
|
|
||||||
this.showVersionStep();
|
this.showVersionStep();
|
||||||
@@ -427,6 +431,9 @@ export class DownloadManager {
|
|||||||
await this.retrieveVersionsForModel(this.modelId, this.source);
|
await this.retrieveVersionsForModel(this.modelId, this.source);
|
||||||
if (this.modelVersionId) {
|
if (this.modelVersionId) {
|
||||||
this.currentVersion = this.versions.find(v => v.id.toString() === this.modelVersionId);
|
this.currentVersion = this.versions.find(v => v.id.toString() === this.modelVersionId);
|
||||||
|
} else {
|
||||||
|
// No explicit version id → default to the latest version (Civitai returns newest first)
|
||||||
|
this.currentVersion = this.versions[0];
|
||||||
}
|
}
|
||||||
this.showVersionStep();
|
this.showVersionStep();
|
||||||
} catch (error) {
|
} catch (error) {
|
||||||
@@ -747,7 +754,7 @@ export class DownloadManager {
|
|||||||
this.selectedFile?.id, this.selectedFile?.name, this.selectedFile?.type, this.selectedFile?.metadata);
|
this.selectedFile?.id, this.selectedFile?.name, this.selectedFile?.type, this.selectedFile?.metadata);
|
||||||
|
|
||||||
document.getElementById('fileSelectionStep').style.display = 'none';
|
document.getElementById('fileSelectionStep').style.display = 'none';
|
||||||
document.getElementById('locationStep').style.display = 'block';
|
document.getElementById('downloadLocationStep').style.display = 'block';
|
||||||
this.proceedToLocationContent();
|
this.proceedToLocationContent();
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -781,7 +788,7 @@ export class DownloadManager {
|
|||||||
}
|
}
|
||||||
|
|
||||||
document.querySelectorAll('.download-step').forEach(step => step.style.display = 'none');
|
document.querySelectorAll('.download-step').forEach(step => step.style.display = 'none');
|
||||||
document.getElementById('locationStep').style.display = 'block';
|
document.getElementById('downloadLocationStep').style.display = 'block';
|
||||||
await this.proceedToLocationContent();
|
await this.proceedToLocationContent();
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -878,6 +885,26 @@ export class DownloadManager {
|
|||||||
this.updateTargetPath();
|
this.updateTargetPath();
|
||||||
}
|
}
|
||||||
|
|
||||||
|
/**
|
||||||
|
* Synthesize a clickable URL for a single-download failure entry.
|
||||||
|
* Single downloads have no pasted URL, so the modal link is derived from
|
||||||
|
* the model/version ids (CivitAI) or the HF repo/file (HuggingFace).
|
||||||
|
*/
|
||||||
|
_buildSingleItemUrl({ modelId, versionId, source, repo = null, filename = null }) {
|
||||||
|
if (source === 'huggingface' && repo) {
|
||||||
|
const base = `https://huggingface.co/${encodeURI(repo)}`;
|
||||||
|
return filename ? `${base}/blob/${encodeURI('main')}/${encodeURI(filename)}` : base;
|
||||||
|
}
|
||||||
|
if (modelId) {
|
||||||
|
return buildCivitaiUrl({
|
||||||
|
modelId,
|
||||||
|
versionId,
|
||||||
|
host: normalizeCivitaiPageHost(state?.global?.settings?.civitai_host),
|
||||||
|
});
|
||||||
|
}
|
||||||
|
return null;
|
||||||
|
}
|
||||||
|
|
||||||
async executeDownloadWithProgress({
|
async executeDownloadWithProgress({
|
||||||
modelId,
|
modelId,
|
||||||
versionId,
|
versionId,
|
||||||
@@ -896,6 +923,7 @@ export class DownloadManager {
|
|||||||
}
|
}
|
||||||
|
|
||||||
const displayName = versionName || `#${versionId}`;
|
const displayName = versionName || `#${versionId}`;
|
||||||
|
const retryParams = { modelId, versionId, versionName, modelRoot, targetFolder, useDefaultPaths, source, fileParams, closeModal: false };
|
||||||
let ws = null;
|
let ws = null;
|
||||||
let updateProgress = () => { };
|
let updateProgress = () => { };
|
||||||
let cancelled = false;
|
let cancelled = false;
|
||||||
@@ -984,6 +1012,26 @@ export class DownloadManager {
|
|||||||
return true;
|
return true;
|
||||||
}
|
}
|
||||||
|
|
||||||
|
if (!response?.success) {
|
||||||
|
this.loadingManager.setStatus(translate('modals.download.status.finalizing'));
|
||||||
|
showDownloadBatchSummary({
|
||||||
|
total: 1,
|
||||||
|
completed: 0,
|
||||||
|
failedItems: [{
|
||||||
|
item: {
|
||||||
|
modelId,
|
||||||
|
versionId,
|
||||||
|
source,
|
||||||
|
url: this._buildSingleItemUrl({ modelId, versionId, source }),
|
||||||
|
},
|
||||||
|
error: response?.error || 'Unknown error',
|
||||||
|
name: displayName,
|
||||||
|
}],
|
||||||
|
onRetry: () => this.executeDownloadWithProgress(retryParams),
|
||||||
|
});
|
||||||
|
return false;
|
||||||
|
}
|
||||||
|
|
||||||
showToast('toast.loras.downloadCompleted', {}, 'success');
|
showToast('toast.loras.downloadCompleted', {}, 'success');
|
||||||
|
|
||||||
if (closeModal) {
|
if (closeModal) {
|
||||||
@@ -1018,7 +1066,21 @@ export class DownloadManager {
|
|||||||
console.log('Download cancelled by user:', downloadId);
|
console.log('Download cancelled by user:', downloadId);
|
||||||
} else {
|
} else {
|
||||||
console.error('Failed to download model version:', error);
|
console.error('Failed to download model version:', error);
|
||||||
showToast('toast.downloads.downloadError', { message: error?.message }, 'error');
|
showDownloadBatchSummary({
|
||||||
|
total: 1,
|
||||||
|
completed: 0,
|
||||||
|
failedItems: [{
|
||||||
|
item: {
|
||||||
|
modelId,
|
||||||
|
versionId,
|
||||||
|
source,
|
||||||
|
url: this._buildSingleItemUrl({ modelId, versionId, source }),
|
||||||
|
},
|
||||||
|
error: error?.message || 'Unknown error',
|
||||||
|
name: displayName,
|
||||||
|
}],
|
||||||
|
onRetry: () => this.executeDownloadWithProgress(retryParams),
|
||||||
|
});
|
||||||
}
|
}
|
||||||
return false;
|
return false;
|
||||||
} finally {
|
} finally {
|
||||||
@@ -1033,14 +1095,16 @@ export class DownloadManager {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
async _downloadHfSingle({ modelRoot, targetFolder, useDefaultPaths }) {
|
async _downloadHfSingle({ modelRoot, targetFolder, useDefaultPaths, files = null }) {
|
||||||
modalManager.closeModal('downloadModal');
|
modalManager.closeModal('downloadModal');
|
||||||
this.loadingManager.restoreProgressBar();
|
this.loadingManager.restoreProgressBar();
|
||||||
const totalFiles = this.hfSelectedFiles.length;
|
const filesToDownload = files || this.hfSelectedFiles;
|
||||||
|
const totalFiles = filesToDownload.length;
|
||||||
const updateProgress = this.loadingManager.showDownloadProgress(totalFiles);
|
const updateProgress = this.loadingManager.showDownloadProgress(totalFiles);
|
||||||
|
|
||||||
let cancelled = false;
|
let cancelled = false;
|
||||||
let currentDownloadId = null;
|
let currentDownloadId = null;
|
||||||
|
const failedFiles = [];
|
||||||
|
|
||||||
this.loadingManager.showCancelButton(async () => {
|
this.loadingManager.showCancelButton(async () => {
|
||||||
if (cancelled) return;
|
if (cancelled) return;
|
||||||
@@ -1059,7 +1123,7 @@ export class DownloadManager {
|
|||||||
for (let i = 0; i < totalFiles; i++) {
|
for (let i = 0; i < totalFiles; i++) {
|
||||||
if (cancelled) break;
|
if (cancelled) break;
|
||||||
|
|
||||||
const filename = this.hfSelectedFiles[i];
|
const filename = filesToDownload[i];
|
||||||
updateProgress(0, completedDownloads, filename);
|
updateProgress(0, completedDownloads, filename);
|
||||||
this.loadingManager.setStatus(`Downloading ${filename}...`);
|
this.loadingManager.setStatus(`Downloading ${filename}...`);
|
||||||
|
|
||||||
@@ -1105,6 +1169,31 @@ export class DownloadManager {
|
|||||||
if (response?.success) {
|
if (response?.success) {
|
||||||
completedDownloads++;
|
completedDownloads++;
|
||||||
updateProgress(100, completedDownloads, filename);
|
updateProgress(100, completedDownloads, filename);
|
||||||
|
} else {
|
||||||
|
failedFiles.push({
|
||||||
|
item: {
|
||||||
|
source: 'huggingface',
|
||||||
|
repo: this.hfRepoId,
|
||||||
|
filename,
|
||||||
|
url: this._buildSingleItemUrl({ source: 'huggingface', repo: this.hfRepoId, filename }),
|
||||||
|
},
|
||||||
|
error: response?.error || 'Unknown error',
|
||||||
|
name: filename,
|
||||||
|
});
|
||||||
|
}
|
||||||
|
} catch (err) {
|
||||||
|
if (!cancelled) {
|
||||||
|
console.error(`Failed to download HF file ${filename}:`, err);
|
||||||
|
failedFiles.push({
|
||||||
|
item: {
|
||||||
|
source: 'huggingface',
|
||||||
|
repo: this.hfRepoId,
|
||||||
|
filename,
|
||||||
|
url: this._buildSingleItemUrl({ source: 'huggingface', repo: this.hfRepoId, filename }),
|
||||||
|
},
|
||||||
|
error: err?.message || 'Unknown error',
|
||||||
|
name: filename,
|
||||||
|
});
|
||||||
}
|
}
|
||||||
} finally {
|
} finally {
|
||||||
ws.close();
|
ws.close();
|
||||||
@@ -1114,11 +1203,27 @@ export class DownloadManager {
|
|||||||
if (cancelled) {
|
if (cancelled) {
|
||||||
showToast('toast.downloads.downloadStopped', {}, 'info',
|
showToast('toast.downloads.downloadStopped', {}, 'info',
|
||||||
`Download cancelled. ${completedDownloads} item(s) completed.`);
|
`Download cancelled. ${completedDownloads} item(s) completed.`);
|
||||||
} else {
|
await resetAndReload(true);
|
||||||
showToast('toast.loras.downloadCompleted', {}, 'success');
|
return true;
|
||||||
}
|
}
|
||||||
|
if (failedFiles.length === 0) {
|
||||||
|
showToast('toast.loras.downloadCompleted', {}, 'success');
|
||||||
|
await resetAndReload(true);
|
||||||
|
return true;
|
||||||
|
}
|
||||||
|
showDownloadBatchSummary({
|
||||||
|
total: totalFiles,
|
||||||
|
completed: completedDownloads,
|
||||||
|
failedItems: failedFiles,
|
||||||
|
onRetry: () => this._downloadHfSingle({
|
||||||
|
modelRoot,
|
||||||
|
targetFolder,
|
||||||
|
useDefaultPaths,
|
||||||
|
files: failedFiles.map((f) => f.item.filename),
|
||||||
|
}),
|
||||||
|
});
|
||||||
await resetAndReload(true);
|
await resetAndReload(true);
|
||||||
return true;
|
return false;
|
||||||
} catch (error) {
|
} catch (error) {
|
||||||
if (!cancelled) {
|
if (!cancelled) {
|
||||||
console.error('Failed to download HF model:', error);
|
console.error('Failed to download HF model:', error);
|
||||||
@@ -1458,7 +1563,7 @@ export class DownloadManager {
|
|||||||
}
|
}
|
||||||
|
|
||||||
backToVersions() {
|
backToVersions() {
|
||||||
document.getElementById('locationStep').style.display = 'none';
|
document.getElementById('downloadLocationStep').style.display = 'none';
|
||||||
if (this.isBatchMode) {
|
if (this.isBatchMode) {
|
||||||
document.getElementById('batchPreviewStep').style.display = 'block';
|
document.getElementById('batchPreviewStep').style.display = 'block';
|
||||||
} else {
|
} else {
|
||||||
@@ -1548,6 +1653,10 @@ export class DownloadManager {
|
|||||||
|
|
||||||
modalManager.closeModal('downloadModal');
|
modalManager.closeModal('downloadModal');
|
||||||
|
|
||||||
|
return this.executeBatchDownload(downloadItems, { modelRoot, targetFolder, useDefaultPaths });
|
||||||
|
}
|
||||||
|
|
||||||
|
async executeBatchDownload(downloadItems, { modelRoot, targetFolder, useDefaultPaths }) {
|
||||||
const batchDownloadId = Date.now().toString();
|
const batchDownloadId = Date.now().toString();
|
||||||
const wsProtocol = window.location.protocol === 'https:' ? 'wss://' : 'ws://';
|
const wsProtocol = window.location.protocol === 'https:' ? 'wss://' : 'ws://';
|
||||||
const ws = new WebSocket(`${wsProtocol}${window.location.host}/ws/download-progress?id=${batchDownloadId}`);
|
const ws = new WebSocket(`${wsProtocol}${window.location.host}/ws/download-progress?id=${batchDownloadId}`);
|
||||||
@@ -1558,6 +1667,7 @@ export class DownloadManager {
|
|||||||
let completedDownloads = 0;
|
let completedDownloads = 0;
|
||||||
let failedDownloads = 0;
|
let failedDownloads = 0;
|
||||||
let cancelled = false;
|
let cancelled = false;
|
||||||
|
const failedItems = [];
|
||||||
|
|
||||||
loadingManager.showCancelButton(async () => {
|
loadingManager.showCancelButton(async () => {
|
||||||
if (cancelled) return;
|
if (cancelled) return;
|
||||||
@@ -1658,6 +1768,7 @@ export class DownloadManager {
|
|||||||
|
|
||||||
if (!response.success) {
|
if (!response.success) {
|
||||||
failedDownloads++;
|
failedDownloads++;
|
||||||
|
failedItems.push({ item, error: response.error || 'Unknown error', name });
|
||||||
} else {
|
} else {
|
||||||
completedDownloads++;
|
completedDownloads++;
|
||||||
updateProgress(100, completedDownloads, '');
|
updateProgress(100, completedDownloads, '');
|
||||||
@@ -1666,6 +1777,7 @@ export class DownloadManager {
|
|||||||
if (!cancelled) {
|
if (!cancelled) {
|
||||||
console.error(`Failed to download ${name}:`, err);
|
console.error(`Failed to download ${name}:`, err);
|
||||||
failedDownloads++;
|
failedDownloads++;
|
||||||
|
failedItems.push({ item, error: err?.message || 'Unknown error', name });
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -1679,10 +1791,15 @@ export class DownloadManager {
|
|||||||
} else if (failedDownloads === 0) {
|
} else if (failedDownloads === 0) {
|
||||||
showToast('toast.loras.allDownloadSuccessful', { count: completedDownloads }, 'success');
|
showToast('toast.loras.allDownloadSuccessful', { count: completedDownloads }, 'success');
|
||||||
} else {
|
} else {
|
||||||
showToast('toast.loras.downloadPartialSuccess', {
|
showDownloadBatchSummary({
|
||||||
completed: completedDownloads,
|
|
||||||
total: downloadItems.length,
|
total: downloadItems.length,
|
||||||
}, 'warning');
|
completed: completedDownloads,
|
||||||
|
failedItems,
|
||||||
|
onRetry: (failed) => this.executeBatchDownload(
|
||||||
|
failed.map((f) => f.item),
|
||||||
|
{ modelRoot, targetFolder, useDefaultPaths }
|
||||||
|
),
|
||||||
|
});
|
||||||
}
|
}
|
||||||
|
|
||||||
await resetAndReload(true);
|
await resetAndReload(true);
|
||||||
|
|||||||
@@ -6,7 +6,6 @@ import { state } from '../state/index.js';
|
|||||||
// Constants for preset management
|
// Constants for preset management
|
||||||
const PRESETS_STORAGE_VERSION = 'v1';
|
const PRESETS_STORAGE_VERSION = 'v1';
|
||||||
const MAX_PRESET_NAME_LENGTH = 30;
|
const MAX_PRESET_NAME_LENGTH = 30;
|
||||||
const MAX_PRESETS_COUNT = 10;
|
|
||||||
|
|
||||||
// Marker for when wildcard patterns resolve to no matches
|
// Marker for when wildcard patterns resolve to no matches
|
||||||
// This ensures we return empty results instead of all models
|
// This ensures we return empty results instead of all models
|
||||||
@@ -362,15 +361,6 @@ export class FilterPresetManager {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
if (presets.length >= MAX_PRESETS_COUNT) {
|
|
||||||
showToast(
|
|
||||||
translate('toast.error.maxPresetsReached', { max: MAX_PRESETS_COUNT }, `Maximum ${MAX_PRESETS_COUNT} presets allowed. Delete one to add more.`),
|
|
||||||
{},
|
|
||||||
'error'
|
|
||||||
);
|
|
||||||
return false;
|
|
||||||
}
|
|
||||||
|
|
||||||
const preset = {
|
const preset = {
|
||||||
name: trimmedName,
|
name: trimmedName,
|
||||||
filters: this.filterManager.cloneFilters(),
|
filters: this.filterManager.cloneFilters(),
|
||||||
@@ -613,17 +603,6 @@ export class FilterPresetManager {
|
|||||||
return;
|
return;
|
||||||
}
|
}
|
||||||
|
|
||||||
// Check max presets limit before showing input
|
|
||||||
const presets = this.loadPresets();
|
|
||||||
if (presets.length >= MAX_PRESETS_COUNT) {
|
|
||||||
showToast(
|
|
||||||
translate('toast.error.maxPresetsReached', { max: MAX_PRESETS_COUNT }, `Maximum ${MAX_PRESETS_COUNT} presets allowed. Delete one to add more.`),
|
|
||||||
{},
|
|
||||||
'error'
|
|
||||||
);
|
|
||||||
return;
|
|
||||||
}
|
|
||||||
|
|
||||||
this.isInlineNamingActive = true;
|
this.isInlineNamingActive = true;
|
||||||
|
|
||||||
const presetsContainer = document.getElementById('filterPresets');
|
const presetsContainer = document.getElementById('filterPresets');
|
||||||
|
|||||||
@@ -1017,6 +1017,12 @@ export class SettingsManager {
|
|||||||
displayDensitySelect.value = state.global.settings.display_density || 'default';
|
displayDensitySelect.value = state.global.settings.display_density || 'default';
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// Set recipes layout setting
|
||||||
|
const recipesLayoutSelect = document.getElementById('recipesLayout');
|
||||||
|
if (recipesLayoutSelect) {
|
||||||
|
recipesLayoutSelect.value = state.global.settings.recipes_layout || 'grid';
|
||||||
|
}
|
||||||
|
|
||||||
// Set card info display setting
|
// Set card info display setting
|
||||||
const cardInfoDisplaySelect = document.getElementById('cardInfoDisplay');
|
const cardInfoDisplaySelect = document.getElementById('cardInfoDisplay');
|
||||||
if (cardInfoDisplaySelect) {
|
if (cardInfoDisplaySelect) {
|
||||||
@@ -2288,6 +2294,13 @@ export class SettingsManager {
|
|||||||
// Apply frontend settings immediately
|
// Apply frontend settings immediately
|
||||||
this.applyFrontendSettings();
|
this.applyFrontendSettings();
|
||||||
|
|
||||||
|
// Dispatch layout change event; the scroller instance is about to be rebuilt,
|
||||||
|
// so calculateLayout() must NOT run on the old instance here
|
||||||
|
if (settingKey === 'recipes_layout') {
|
||||||
|
window.dispatchEvent(new CustomEvent('lm:recipes-layout-changed'));
|
||||||
|
return;
|
||||||
|
}
|
||||||
|
|
||||||
// Recalculate layout when display density changes
|
// Recalculate layout when display density changes
|
||||||
if (settingKey === 'display_density' && state.virtualScroller) {
|
if (settingKey === 'display_density' && state.virtualScroller) {
|
||||||
state.virtualScroller.calculateLayout();
|
state.virtualScroller.calculateLayout();
|
||||||
|
|||||||
@@ -47,11 +47,8 @@ export class ImportStepManager {
|
|||||||
targetStep.offsetHeight;
|
targetStep.offsetHeight;
|
||||||
}
|
}
|
||||||
|
|
||||||
// Scroll modal content to top
|
// Scroll the active step back to top (steps scroll independently of the modal shell)
|
||||||
const modalContent = document.querySelector('#importModal .modal-content');
|
targetStep.scrollTop = 0;
|
||||||
if (modalContent) {
|
|
||||||
modalContent.scrollTop = 0;
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
+15
-1
@@ -7,7 +7,7 @@ import { state, getCurrentPageState } from './state/index.js';
|
|||||||
import { getStorageItem, setStorageItem, getSessionItem, removeSessionItem } from './utils/storageHelpers.js';
|
import { getStorageItem, setStorageItem, getSessionItem, removeSessionItem } from './utils/storageHelpers.js';
|
||||||
import { RecipeContextMenu } from './components/ContextMenu/index.js';
|
import { RecipeContextMenu } from './components/ContextMenu/index.js';
|
||||||
import { DuplicatesManager } from './components/DuplicatesManager.js';
|
import { DuplicatesManager } from './components/DuplicatesManager.js';
|
||||||
import { refreshVirtualScroll } from './utils/infiniteScroll.js';
|
import { refreshVirtualScroll, recreateVirtualScroll } from './utils/infiniteScroll.js';
|
||||||
import { refreshRecipes, RecipeSidebarApiClient } from './api/recipeApi.js';
|
import { refreshRecipes, RecipeSidebarApiClient } from './api/recipeApi.js';
|
||||||
import { sidebarManager } from './components/SidebarManager.js';
|
import { sidebarManager } from './components/SidebarManager.js';
|
||||||
import { initSortDropdown } from './components/controls/SortDropdown.js';
|
import { initSortDropdown } from './components/controls/SortDropdown.js';
|
||||||
@@ -272,6 +272,20 @@ class RecipeManager {
|
|||||||
});
|
});
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// Rebuild the scroller on layout switch; in duplicates mode defer until
|
||||||
|
// exitDuplicateMode re-enables the scroller (direct recreation would dispose
|
||||||
|
// the old instance while initializeVirtualScroll skips duplicates mode)
|
||||||
|
window.addEventListener('lm:recipes-layout-changed', () => {
|
||||||
|
const pageState = getCurrentPageState();
|
||||||
|
if (pageState.duplicatesMode) {
|
||||||
|
state.pendingLayoutRecreate = true;
|
||||||
|
return;
|
||||||
|
}
|
||||||
|
if (typeof recreateVirtualScroll === 'function') {
|
||||||
|
recreateVirtualScroll('recipes');
|
||||||
|
}
|
||||||
|
});
|
||||||
|
|
||||||
// Initialize dropdown functionality for refresh button
|
// Initialize dropdown functionality for refresh button
|
||||||
this.initDropdowns();
|
this.initDropdowns();
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -38,6 +38,7 @@ const DEFAULT_SETTINGS_BASE = Object.freeze({
|
|||||||
card_blur_amount: 8,
|
card_blur_amount: 8,
|
||||||
autoplay_on_hover: false,
|
autoplay_on_hover: false,
|
||||||
display_density: 'default',
|
display_density: 'default',
|
||||||
|
recipes_layout: 'grid',
|
||||||
card_info_display: 'always',
|
card_info_display: 'always',
|
||||||
model_name_display: 'model_name',
|
model_name_display: 'model_name',
|
||||||
lora_syntax_format: 'legacy',
|
lora_syntax_format: 'legacy',
|
||||||
|
|||||||
File diff suppressed because it is too large
Load Diff
@@ -483,10 +483,18 @@ export class VirtualScroller {
|
|||||||
element.style.width = `${this.itemWidth}px`;
|
element.style.width = `${this.itemWidth}px`;
|
||||||
element.style.height = `${this.itemHeight}px`;
|
element.style.height = `${this.itemHeight}px`;
|
||||||
|
|
||||||
// Remove max-width constraint from model-card to allow dynamic sizing
|
// Remove max-width/min-width constraints from the model-card to allow
|
||||||
const modelCard = element.querySelector('.model-card');
|
// dynamic sizing. The card is either the item element itself (e.g.
|
||||||
|
// RecipeCard returns the .model-card root) or a descendant (ModelCard).
|
||||||
|
// Without this, the CSS min-width of 200px forces compact-density cards
|
||||||
|
// wider than their allocated column, and the last column gets clipped
|
||||||
|
// by the grid's overflow-x: hidden.
|
||||||
|
const modelCard = element.classList.contains('model-card')
|
||||||
|
? element
|
||||||
|
: element.querySelector('.model-card');
|
||||||
if (modelCard) {
|
if (modelCard) {
|
||||||
modelCard.style.maxWidth = 'none';
|
modelCard.style.maxWidth = 'none';
|
||||||
|
modelCard.style.minWidth = '0';
|
||||||
}
|
}
|
||||||
|
|
||||||
return element;
|
return element;
|
||||||
|
|||||||
@@ -1,5 +1,6 @@
|
|||||||
import { state, getCurrentPageState } from '../state/index.js';
|
import { state, getCurrentPageState } from '../state/index.js';
|
||||||
import { VirtualScroller } from './VirtualScroller.js';
|
import { VirtualScroller } from './VirtualScroller.js';
|
||||||
|
import { MasonryScroller } from './MasonryScroller.js';
|
||||||
import { createModelCard, setupModelCardEventDelegation } from '../components/shared/ModelCard.js';
|
import { createModelCard, setupModelCardEventDelegation } from '../components/shared/ModelCard.js';
|
||||||
import { getModelApiClient } from '../api/modelApiFactory.js';
|
import { getModelApiClient } from '../api/modelApiFactory.js';
|
||||||
import { showToast } from './uiHelpers.js';
|
import { showToast } from './uiHelpers.js';
|
||||||
@@ -141,8 +142,14 @@ async function initializeVirtualScroll(pageType) {
|
|||||||
throw new Error(`Required components not available for ${pageType} page`);
|
throw new Error(`Required components not available for ${pageType} page`);
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// Masonry applies to the recipes page only; the read pattern mirrors
|
||||||
|
// VirtualScroller.js (display_density).
|
||||||
|
const useMasonry = pageType === 'recipes'
|
||||||
|
&& (state.global.settings?.recipes_layout ?? 'grid') === 'masonry';
|
||||||
|
const ScrollerClass = useMasonry ? MasonryScroller : VirtualScroller;
|
||||||
|
|
||||||
// Initialize virtual scroller with renamed container elements
|
// Initialize virtual scroller with renamed container elements
|
||||||
state.virtualScroller = new VirtualScroller({
|
state.virtualScroller = new ScrollerClass({
|
||||||
gridElement: grid,
|
gridElement: grid,
|
||||||
containerElement: gridContainer,
|
containerElement: gridContainer,
|
||||||
scrollContainer: scrollContainer,
|
scrollContainer: scrollContainer,
|
||||||
@@ -250,3 +257,19 @@ export async function refreshVirtualScroll(options = {}) {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// Rebuild the virtual scroller from scratch (used for layout switching).
|
||||||
|
// refreshVirtualScroll only resets the existing instance, so it cannot swap
|
||||||
|
// the scroller class. Keyboard navigation must be cleaned up first:
|
||||||
|
// setupKeyboardNavigation appends a new document keydown listener on every
|
||||||
|
// initializeVirtualScroll call and never removes the previous one.
|
||||||
|
export async function recreateVirtualScroll(pageType) {
|
||||||
|
cleanupKeyboardNavigation();
|
||||||
|
|
||||||
|
if (state.virtualScroller) {
|
||||||
|
state.virtualScroller.dispose();
|
||||||
|
state.virtualScroller = null;
|
||||||
|
}
|
||||||
|
|
||||||
|
await initializeVirtualScroll(pageType);
|
||||||
|
}
|
||||||
|
|||||||
@@ -356,6 +356,8 @@ export function updatePanelPositions() {
|
|||||||
|
|
||||||
if (filterPanel) {
|
if (filterPanel) {
|
||||||
filterPanel.style.top = `${topPosition}px`;
|
filterPanel.style.top = `${topPosition}px`;
|
||||||
|
// Clamp panel height to the viewport below the header
|
||||||
|
filterPanel.style.maxHeight = `calc(100vh - ${topPosition + 10}px)`;
|
||||||
}
|
}
|
||||||
|
|
||||||
// Adjust panel horizontal position based on the search container
|
// Adjust panel horizontal position based on the search container
|
||||||
|
|||||||
@@ -32,7 +32,12 @@
|
|||||||
<div class="context-menu-separator menu-section-break"></div>
|
<div class="context-menu-separator menu-section-break"></div>
|
||||||
<!-- Media / Preview -->
|
<!-- Media / Preview -->
|
||||||
<div class="context-menu-item" data-action="preview"><i class="fas fa-folder-open"></i> {{ t('loras.contextMenu.openExamples') }}</div>
|
<div class="context-menu-item" data-action="preview"><i class="fas fa-folder-open"></i> {{ t('loras.contextMenu.openExamples') }}</div>
|
||||||
<div class="context-menu-item" data-action="download-examples"><i class="fas fa-download"></i> {{ t('loras.contextMenu.downloadExamples') }}</div>
|
<div class="context-menu-item has-submenu" data-has-submenu="download-examples"><i class="fas fa-download"></i> {{ t('loras.contextMenu.downloadExamples') }} <i class="fas fa-chevron-right submenu-arrow"></i>
|
||||||
|
<div class="context-submenu">
|
||||||
|
<div class="context-menu-item" data-action="download-examples"><i class="fas fa-download"></i> {{ t('loras.contextMenu.downloadMissingExamples') }}</div>
|
||||||
|
<div class="context-menu-item" data-action="download-examples-force"><i class="fas fa-redo-alt"></i> {{ t('loras.contextMenu.reprocessExamples') }}</div>
|
||||||
|
</div>
|
||||||
|
</div>
|
||||||
<div class="context-menu-item" data-action="replace-preview"><i class="fas fa-image"></i> {{ t('loras.contextMenu.replacePreview') }}</div>
|
<div class="context-menu-item" data-action="replace-preview"><i class="fas fa-image"></i> {{ t('loras.contextMenu.replacePreview') }}</div>
|
||||||
<div class="context-menu-separator menu-section-break"></div>
|
<div class="context-menu-separator menu-section-break"></div>
|
||||||
<!-- Attributes -->
|
<!-- Attributes -->
|
||||||
|
|||||||
@@ -44,8 +44,18 @@
|
|||||||
<div class="context-menu-item" data-action="preview">
|
<div class="context-menu-item" data-action="preview">
|
||||||
<i class="fas fa-folder-open"></i> <span>{{ t('loras.contextMenu.openExamples') }}</span>
|
<i class="fas fa-folder-open"></i> <span>{{ t('loras.contextMenu.openExamples') }}</span>
|
||||||
</div>
|
</div>
|
||||||
<div class="context-menu-item" data-action="download-examples">
|
<div class="context-menu-item has-submenu" data-has-submenu="download-examples">
|
||||||
<i class="fas fa-download"></i> <span>{{ t('loras.contextMenu.downloadExamples') }}</span>
|
<i class="fas fa-download"></i>
|
||||||
|
<span>{{ t('loras.contextMenu.downloadExamples') }}</span>
|
||||||
|
<i class="fas fa-chevron-right submenu-arrow"></i>
|
||||||
|
<div class="context-submenu">
|
||||||
|
<div class="context-menu-item" data-action="download-examples">
|
||||||
|
<i class="fas fa-download"></i> <span>{{ t('loras.contextMenu.downloadMissingExamples') }}</span>
|
||||||
|
</div>
|
||||||
|
<div class="context-menu-item" data-action="download-examples-force">
|
||||||
|
<i class="fas fa-redo-alt"></i> <span>{{ t('loras.contextMenu.reprocessExamples') }}</span>
|
||||||
|
</div>
|
||||||
|
</div>
|
||||||
</div>
|
</div>
|
||||||
<div class="context-menu-item" data-action="replace-preview">
|
<div class="context-menu-item" data-action="replace-preview">
|
||||||
<i class="fas fa-image"></i> <span>{{ t('loras.contextMenu.replacePreview') }}</span>
|
<i class="fas fa-image"></i> <span>{{ t('loras.contextMenu.replacePreview') }}</span>
|
||||||
@@ -136,8 +146,18 @@
|
|||||||
</div>
|
</div>
|
||||||
<div class="context-menu-section" data-section="download">
|
<div class="context-menu-section" data-section="download">
|
||||||
<div class="context-menu-section-header">{{ t('loras.bulkOperations.sections.download') }}</div>
|
<div class="context-menu-section-header">{{ t('loras.bulkOperations.sections.download') }}</div>
|
||||||
<div class="context-menu-item" data-action="download-example-images">
|
<div class="context-menu-item has-submenu" data-has-submenu="download-example-images">
|
||||||
<i class="fas fa-download"></i> <span>{{ t('loras.bulkOperations.downloadExamples') }}</span>
|
<i class="fas fa-download"></i>
|
||||||
|
<span>{{ t('loras.bulkOperations.downloadExamples') }}</span>
|
||||||
|
<i class="fas fa-chevron-right submenu-arrow"></i>
|
||||||
|
<div class="context-submenu">
|
||||||
|
<div class="context-menu-item" data-action="download-missing-example-images">
|
||||||
|
<i class="fas fa-download"></i> <span>{{ t('loras.bulkOperations.downloadMissingExamples') }}</span>
|
||||||
|
</div>
|
||||||
|
<div class="context-menu-item" data-action="download-example-images">
|
||||||
|
<i class="fas fa-redo-alt"></i> <span>{{ t('loras.bulkOperations.reprocessExamples') }}</span>
|
||||||
|
</div>
|
||||||
|
</div>
|
||||||
</div>
|
</div>
|
||||||
<div class="context-menu-item" data-action="download-missing-loras">
|
<div class="context-menu-item" data-action="download-missing-loras">
|
||||||
<i class="fas fa-download"></i> <span>{{ t('loras.bulkOperations.downloadMissingLoras') }}</span>
|
<i class="fas fa-download"></i> <span>{{ t('loras.bulkOperations.downloadMissingLoras') }}</span>
|
||||||
|
|||||||
@@ -48,6 +48,11 @@
|
|||||||
<option value="versions_count:asc">{{ t('loras.controls.sort.versionsCountAsc', default='Fewest versions first') }}</option>
|
<option value="versions_count:asc">{{ t('loras.controls.sort.versionsCountAsc', default='Fewest versions first') }}</option>
|
||||||
</optgroup>
|
</optgroup>
|
||||||
{% endif %}
|
{% endif %}
|
||||||
|
{% if page_id != 'recipes' %}
|
||||||
|
<optgroup label="{{ t('loras.controls.sort.random', default='Random') }}">
|
||||||
|
<option value="random">{{ t('loras.controls.sort.randomAction', default='Randomize (shuffle)') }}</option>
|
||||||
|
</optgroup>
|
||||||
|
{% endif %}
|
||||||
{% if page_id == 'recipes' %}
|
{% if page_id == 'recipes' %}
|
||||||
<optgroup label="{{ t('recipes.controls.sort.lorasCount') }}">
|
<optgroup label="{{ t('recipes.controls.sort.lorasCount') }}">
|
||||||
<option value="loras_count:desc">{{ t('recipes.controls.sort.lorasCountDesc') }}</option>
|
<option value="loras_count:desc">{{ t('recipes.controls.sort.lorasCountDesc') }}</option>
|
||||||
|
|||||||
@@ -60,7 +60,7 @@
|
|||||||
</div>
|
</div>
|
||||||
|
|
||||||
<!-- Step 3: Location Selection -->
|
<!-- Step 3: Location Selection -->
|
||||||
<div class="download-step" id="locationStep" style="display: none;">
|
<div class="download-step" id="downloadLocationStep" style="display: none;">
|
||||||
<div class="location-selection">
|
<div class="location-selection">
|
||||||
<!-- Path preview with inline toggle -->
|
<!-- Path preview with inline toggle -->
|
||||||
<div class="path-preview">
|
<div class="path-preview">
|
||||||
|
|||||||
@@ -202,6 +202,7 @@
|
|||||||
<option value="deepseek">{{ t('settings.aiProvider.providerOptions.deepseek') }}</option>
|
<option value="deepseek">{{ t('settings.aiProvider.providerOptions.deepseek') }}</option>
|
||||||
<option value="groq">{{ t('settings.aiProvider.providerOptions.groq') }}</option>
|
<option value="groq">{{ t('settings.aiProvider.providerOptions.groq') }}</option>
|
||||||
<option value="openrouter">{{ t('settings.aiProvider.providerOptions.openrouter') }}</option>
|
<option value="openrouter">{{ t('settings.aiProvider.providerOptions.openrouter') }}</option>
|
||||||
|
<option value="google">{{ t('settings.aiProvider.providerOptions.google') }}</option>
|
||||||
<option value="opencode-go">{{ t('settings.aiProvider.providerOptions.opencode-go') }}</option>
|
<option value="opencode-go">{{ t('settings.aiProvider.providerOptions.opencode-go') }}</option>
|
||||||
<option value="custom">{{ t('settings.aiProvider.providerOptions.custom') }}</option>
|
<option value="custom">{{ t('settings.aiProvider.providerOptions.custom') }}</option>
|
||||||
</select>
|
</select>
|
||||||
@@ -624,7 +625,24 @@
|
|||||||
<span class="warning-text">{{ t('settings.layoutSettings.displayDensityWarning') }}</span>
|
<span class="warning-text">{{ t('settings.layoutSettings.displayDensityWarning') }}</span>
|
||||||
</div>
|
</div>
|
||||||
</div>
|
</div>
|
||||||
|
|
||||||
|
<div class="setting-item">
|
||||||
|
<div class="setting-row">
|
||||||
|
<div class="setting-info">
|
||||||
|
<label for="recipesLayout">
|
||||||
|
{{ t('settings.layoutSettings.recipesLayout') }}
|
||||||
|
<i class="fas fa-info-circle info-icon" data-tooltip="{{ t('settings.layoutSettings.recipesLayoutHelp') }}"></i>
|
||||||
|
</label>
|
||||||
|
</div>
|
||||||
|
<div class="setting-control select-control">
|
||||||
|
<select id="recipesLayout" onchange="settingsManager.saveSelectSetting('recipesLayout', 'recipes_layout')">
|
||||||
|
<option value="grid">{{ t('settings.layoutSettings.recipesLayoutOptions.grid') }}</option>
|
||||||
|
<option value="masonry">{{ t('settings.layoutSettings.recipesLayoutOptions.masonry') }}</option>
|
||||||
|
</select>
|
||||||
|
</div>
|
||||||
|
</div>
|
||||||
|
</div>
|
||||||
|
|
||||||
<div class="setting-item">
|
<div class="setting-item">
|
||||||
<div class="setting-row">
|
<div class="setting-row">
|
||||||
<div class="setting-info">
|
<div class="setting-info">
|
||||||
|
|||||||
@@ -32,7 +32,12 @@
|
|||||||
<div class="context-menu-separator menu-section-break"></div>
|
<div class="context-menu-separator menu-section-break"></div>
|
||||||
<!-- Media / Preview -->
|
<!-- Media / Preview -->
|
||||||
<div class="context-menu-item" data-action="preview"><i class="fas fa-folder-open"></i> {{ t('loras.contextMenu.openExamples') }}</div>
|
<div class="context-menu-item" data-action="preview"><i class="fas fa-folder-open"></i> {{ t('loras.contextMenu.openExamples') }}</div>
|
||||||
<div class="context-menu-item" data-action="download-examples"><i class="fas fa-download"></i> {{ t('loras.contextMenu.downloadExamples') }}</div>
|
<div class="context-menu-item has-submenu" data-has-submenu="download-examples"><i class="fas fa-download"></i> {{ t('loras.contextMenu.downloadExamples') }} <i class="fas fa-chevron-right submenu-arrow"></i>
|
||||||
|
<div class="context-submenu">
|
||||||
|
<div class="context-menu-item" data-action="download-examples"><i class="fas fa-download"></i> {{ t('loras.contextMenu.downloadMissingExamples') }}</div>
|
||||||
|
<div class="context-menu-item" data-action="download-examples-force"><i class="fas fa-redo-alt"></i> {{ t('loras.contextMenu.reprocessExamples') }}</div>
|
||||||
|
</div>
|
||||||
|
</div>
|
||||||
<div class="context-menu-item" data-action="replace-preview"><i class="fas fa-image"></i> {{ t('loras.contextMenu.replacePreview') }}</div>
|
<div class="context-menu-item" data-action="replace-preview"><i class="fas fa-image"></i> {{ t('loras.contextMenu.replacePreview') }}</div>
|
||||||
<div class="context-menu-separator menu-section-break"></div>
|
<div class="context-menu-separator menu-section-break"></div>
|
||||||
<!-- Attributes -->
|
<!-- Attributes -->
|
||||||
|
|||||||
@@ -38,6 +38,7 @@ vi.mock('../../../static/js/state/index.js', () => {
|
|||||||
vi.mock('../../../static/js/utils/infiniteScroll.js', () => ({
|
vi.mock('../../../static/js/utils/infiniteScroll.js', () => ({
|
||||||
captureScrollPosition: captureScrollPositionMock,
|
captureScrollPosition: captureScrollPositionMock,
|
||||||
restoreScrollPosition: restoreScrollPositionMock,
|
restoreScrollPosition: restoreScrollPositionMock,
|
||||||
|
recreateVirtualScroll: vi.fn(),
|
||||||
}));
|
}));
|
||||||
|
|
||||||
import {
|
import {
|
||||||
|
|||||||
@@ -1666,4 +1666,374 @@ describe('AutoComplete widget interactions', () => {
|
|||||||
// Entire phrase should be replaced with selected tag
|
// Entire phrase should be replaced with selected tag
|
||||||
expect(input.value).toBe('looking_to_the_side,');
|
expect(input.value).toBe('looking_to_the_side,');
|
||||||
});
|
});
|
||||||
|
|
||||||
|
it('shows /af command for loras when active-filters autocomplete is off (default)', async () => {
|
||||||
|
const input = document.createElement('textarea');
|
||||||
|
input.value = '/';
|
||||||
|
input.selectionStart = input.value.length;
|
||||||
|
document.body.append(input);
|
||||||
|
|
||||||
|
caretHelperInstance.getBeforeCursor.mockReturnValue('/');
|
||||||
|
caretHelperInstance.getCursorOffset.mockReturnValue({ left: 15, top: 25 });
|
||||||
|
|
||||||
|
const { AutoComplete } = await import(AUTOCOMPLETE_MODULE);
|
||||||
|
const autoComplete = new AutoComplete(input, 'loras', { showPreview: false, minChars: 1 });
|
||||||
|
|
||||||
|
input.dispatchEvent(new Event('input', { bubbles: true }));
|
||||||
|
|
||||||
|
const commandNames = autoComplete.items.map((item) => item.command);
|
||||||
|
expect(commandNames).toContain('/af');
|
||||||
|
expect(commandNames).not.toContain('/noaf');
|
||||||
|
expect(commandNames).toContain('/activefilters');
|
||||||
|
expect(commandNames).not.toContain('/noactivefilters');
|
||||||
|
});
|
||||||
|
|
||||||
|
it('does not trigger preview for command items when selecting the loras command list', async () => {
|
||||||
|
// Regression: with showPreview enabled (the default for loras widgets), the
|
||||||
|
// auto-selected first command item was passed to showPreviewForItem() as a
|
||||||
|
// relative path, crashing on relativePath.split.
|
||||||
|
const input = document.createElement('textarea');
|
||||||
|
input.value = '/';
|
||||||
|
input.selectionStart = input.value.length;
|
||||||
|
document.body.append(input);
|
||||||
|
|
||||||
|
caretHelperInstance.getBeforeCursor.mockReturnValue('/');
|
||||||
|
caretHelperInstance.getCursorOffset.mockReturnValue({ left: 15, top: 25 });
|
||||||
|
|
||||||
|
const { AutoComplete } = await import(AUTOCOMPLETE_MODULE);
|
||||||
|
const autoComplete = new AutoComplete(input, 'loras', { showPreview: true, minChars: 1 });
|
||||||
|
|
||||||
|
input.dispatchEvent(new Event('input', { bubbles: true }));
|
||||||
|
|
||||||
|
// Allow the async preview tooltip import to resolve
|
||||||
|
await Promise.resolve();
|
||||||
|
await Promise.resolve();
|
||||||
|
|
||||||
|
const commandNames = autoComplete.items.map((item) => item.command);
|
||||||
|
expect(commandNames).toContain('/af');
|
||||||
|
expect(previewTooltipMock.show).not.toHaveBeenCalled();
|
||||||
|
});
|
||||||
|
|
||||||
|
it('shows /noaf command for loras when active-filters autocomplete is on', async () => {
|
||||||
|
settingGetMock.mockImplementation((key) => {
|
||||||
|
if (key === 'loramanager.lora_active_filters_autocomplete') {
|
||||||
|
return true;
|
||||||
|
}
|
||||||
|
return undefined;
|
||||||
|
});
|
||||||
|
|
||||||
|
const input = document.createElement('textarea');
|
||||||
|
input.value = '/';
|
||||||
|
input.selectionStart = input.value.length;
|
||||||
|
document.body.append(input);
|
||||||
|
|
||||||
|
caretHelperInstance.getBeforeCursor.mockReturnValue('/');
|
||||||
|
caretHelperInstance.getCursorOffset.mockReturnValue({ left: 15, top: 25 });
|
||||||
|
|
||||||
|
const { AutoComplete } = await import(AUTOCOMPLETE_MODULE);
|
||||||
|
const autoComplete = new AutoComplete(input, 'loras', { showPreview: false, minChars: 1 });
|
||||||
|
|
||||||
|
input.dispatchEvent(new Event('input', { bubbles: true }));
|
||||||
|
|
||||||
|
const commandNames = autoComplete.items.map((item) => item.command);
|
||||||
|
expect(commandNames).toContain('/noaf');
|
||||||
|
expect(commandNames).not.toContain('/af');
|
||||||
|
expect(commandNames).toContain('/noactivefilters');
|
||||||
|
expect(commandNames).not.toContain('/activefilters');
|
||||||
|
});
|
||||||
|
|
||||||
|
it('toggles the active-filters setting when /activefilters alias is used', async () => {
|
||||||
|
const input = document.createElement('textarea');
|
||||||
|
input.value = '/activefilters';
|
||||||
|
input.selectionStart = input.value.length;
|
||||||
|
input.focus = vi.fn();
|
||||||
|
input.setSelectionRange = vi.fn();
|
||||||
|
document.body.append(input);
|
||||||
|
|
||||||
|
caretHelperInstance.getBeforeCursor.mockReturnValue('/activefilters');
|
||||||
|
caretHelperInstance.getCursorOffset.mockReturnValue({ left: 15, top: 25 });
|
||||||
|
|
||||||
|
const { AutoComplete } = await import(AUTOCOMPLETE_MODULE);
|
||||||
|
const autoComplete = new AutoComplete(input, 'loras', { showPreview: false, minChars: 1 });
|
||||||
|
|
||||||
|
const commandResult = autoComplete._parseCommandInput('/activefilters');
|
||||||
|
expect(commandResult.command).toBeDefined();
|
||||||
|
expect(commandResult.command.type).toBe('toggle_setting');
|
||||||
|
expect(commandResult.command.value).toBe(true);
|
||||||
|
|
||||||
|
await autoComplete._handleToggleSettingCommand(commandResult.command);
|
||||||
|
|
||||||
|
expect(settingSetMock).toHaveBeenCalledWith('loramanager.lora_active_filters_autocomplete', true);
|
||||||
|
});
|
||||||
|
|
||||||
|
it('toggles the active-filters setting when /af is accepted', async () => {
|
||||||
|
const input = document.createElement('textarea');
|
||||||
|
input.value = '/';
|
||||||
|
input.selectionStart = input.value.length;
|
||||||
|
input.focus = vi.fn();
|
||||||
|
input.setSelectionRange = vi.fn();
|
||||||
|
document.body.append(input);
|
||||||
|
|
||||||
|
caretHelperInstance.getBeforeCursor.mockReturnValue('/');
|
||||||
|
caretHelperInstance.getCursorOffset.mockReturnValue({ left: 15, top: 25 });
|
||||||
|
|
||||||
|
const { AutoComplete } = await import(AUTOCOMPLETE_MODULE);
|
||||||
|
const autoComplete = new AutoComplete(input, 'loras', { showPreview: false, minChars: 1 });
|
||||||
|
|
||||||
|
input.dispatchEvent(new Event('input', { bubbles: true }));
|
||||||
|
|
||||||
|
const afItem = autoComplete.items.find((item) => item.command === '/af');
|
||||||
|
expect(afItem).toBeDefined();
|
||||||
|
|
||||||
|
// Simulate the input being cleared after the command is accepted so the
|
||||||
|
// cleared-token input event does not re-trigger command parsing.
|
||||||
|
caretHelperInstance.getBeforeCursor.mockReturnValue('');
|
||||||
|
await autoComplete._handleToggleSettingCommand(afItem);
|
||||||
|
|
||||||
|
expect(settingSetMock).toHaveBeenCalledWith('loramanager.lora_active_filters_autocomplete', true);
|
||||||
|
});
|
||||||
|
|
||||||
|
it('appends active filter params to loras autocomplete requests when enabled', async () => {
|
||||||
|
vi.useFakeTimers();
|
||||||
|
|
||||||
|
settingGetMock.mockImplementation((key) => {
|
||||||
|
if (key === 'loramanager.lora_active_filters_autocomplete') {
|
||||||
|
return true;
|
||||||
|
}
|
||||||
|
return undefined;
|
||||||
|
});
|
||||||
|
|
||||||
|
localStorage.setItem('lora_manager_loras_filters', JSON.stringify({
|
||||||
|
baseModel: ['SD 1.5'],
|
||||||
|
tags: { anime: 'include', nsfw: 'exclude', __no_tags__: 'exclude' },
|
||||||
|
autoTags: { I2V: 'include' },
|
||||||
|
modelTypes: ['standard'],
|
||||||
|
tagLogic: 'all',
|
||||||
|
license: { noCredit: 'include', allowSelling: 'exclude' },
|
||||||
|
}));
|
||||||
|
localStorage.setItem('lora_manager_loras_activeFolder', 'MyLoras');
|
||||||
|
localStorage.setItem('lora_manager_loras_recursiveSearch', 'true');
|
||||||
|
|
||||||
|
fetchApiMock.mockResolvedValue({
|
||||||
|
json: () => Promise.resolve({ success: true, relative_paths: ['models/example.safetensors'] }),
|
||||||
|
});
|
||||||
|
|
||||||
|
caretHelperInstance.getBeforeCursor.mockReturnValue('example');
|
||||||
|
caretHelperInstance.getCursorOffset.mockReturnValue({ left: 15, top: 25 });
|
||||||
|
|
||||||
|
const input = document.createElement('textarea');
|
||||||
|
document.body.append(input);
|
||||||
|
|
||||||
|
const { AutoComplete } = await import(AUTOCOMPLETE_MODULE);
|
||||||
|
new AutoComplete(input, 'loras', { debounceDelay: 0, showPreview: false, minChars: 1 });
|
||||||
|
|
||||||
|
input.value = 'example';
|
||||||
|
input.dispatchEvent(new Event('input', { bubbles: true }));
|
||||||
|
|
||||||
|
await vi.runAllTimersAsync();
|
||||||
|
await Promise.resolve();
|
||||||
|
|
||||||
|
const calledUrl = fetchApiMock.mock.calls[0][0];
|
||||||
|
expect(calledUrl).toContain('/lm/loras/relative-paths?search=example&limit=100');
|
||||||
|
expect(calledUrl).toContain('folder=MyLoras');
|
||||||
|
expect(calledUrl).toContain('recursive=true');
|
||||||
|
expect(calledUrl).toContain('tag_include=anime');
|
||||||
|
expect(calledUrl).toContain('tag_exclude=nsfw');
|
||||||
|
expect(calledUrl).toContain('tag_exclude=__no_tags__');
|
||||||
|
expect(calledUrl).toContain('auto_tag_include=I2V');
|
||||||
|
expect(calledUrl).toContain('tag_logic=all');
|
||||||
|
expect(calledUrl).toContain('credit_required=false');
|
||||||
|
expect(calledUrl).toContain('allow_selling_generated_content=false');
|
||||||
|
const parsed = new URL(calledUrl, 'https://example.com');
|
||||||
|
expect(parsed.searchParams.get('base_model')).toBe('SD 1.5');
|
||||||
|
expect(parsed.searchParams.get('model_type')).toBe('standard');
|
||||||
|
});
|
||||||
|
|
||||||
|
it('keeps the default loras autocomplete URL when active-filters mode is off', async () => {
|
||||||
|
vi.useFakeTimers();
|
||||||
|
|
||||||
|
fetchApiMock.mockResolvedValue({
|
||||||
|
json: () => Promise.resolve({ success: true, relative_paths: ['models/example.safetensors'] }),
|
||||||
|
});
|
||||||
|
|
||||||
|
caretHelperInstance.getBeforeCursor.mockReturnValue('example');
|
||||||
|
caretHelperInstance.getCursorOffset.mockReturnValue({ left: 15, top: 25 });
|
||||||
|
|
||||||
|
const input = document.createElement('textarea');
|
||||||
|
document.body.append(input);
|
||||||
|
|
||||||
|
const { AutoComplete } = await import(AUTOCOMPLETE_MODULE);
|
||||||
|
new AutoComplete(input, 'loras', { debounceDelay: 0, showPreview: false, minChars: 1 });
|
||||||
|
|
||||||
|
input.value = 'example';
|
||||||
|
input.dispatchEvent(new Event('input', { bubbles: true }));
|
||||||
|
|
||||||
|
await vi.runAllTimersAsync();
|
||||||
|
await Promise.resolve();
|
||||||
|
|
||||||
|
expect(fetchApiMock).toHaveBeenCalledWith('/lm/loras/relative-paths?search=example&limit=100');
|
||||||
|
});
|
||||||
|
|
||||||
|
it('sends the filter-pipeline signal even when no filters are stored', async () => {
|
||||||
|
// Regression: with filter mode on but no folder/filters stored, the request
|
||||||
|
// carried no params, so the backend skipped the filter pipeline and global
|
||||||
|
// settings like show_only_sfw diverged from the list endpoint.
|
||||||
|
vi.useFakeTimers();
|
||||||
|
|
||||||
|
settingGetMock.mockImplementation((key) => {
|
||||||
|
if (key === 'loramanager.lora_active_filters_autocomplete') {
|
||||||
|
return true;
|
||||||
|
}
|
||||||
|
return undefined;
|
||||||
|
});
|
||||||
|
|
||||||
|
localStorage.removeItem('lora_manager_loras_filters');
|
||||||
|
localStorage.removeItem('lora_manager_loras_activeFolder');
|
||||||
|
localStorage.removeItem('lora_manager_loras_recursiveSearch');
|
||||||
|
|
||||||
|
fetchApiMock.mockResolvedValue({
|
||||||
|
json: () => Promise.resolve({ success: true, relative_paths: ['models/example.safetensors'] }),
|
||||||
|
});
|
||||||
|
|
||||||
|
caretHelperInstance.getBeforeCursor.mockReturnValue('example');
|
||||||
|
caretHelperInstance.getCursorOffset.mockReturnValue({ left: 15, top: 25 });
|
||||||
|
|
||||||
|
const input = document.createElement('textarea');
|
||||||
|
document.body.append(input);
|
||||||
|
|
||||||
|
const { AutoComplete } = await import(AUTOCOMPLETE_MODULE);
|
||||||
|
new AutoComplete(input, 'loras', { debounceDelay: 0, showPreview: false, minChars: 1 });
|
||||||
|
|
||||||
|
input.value = 'example';
|
||||||
|
input.dispatchEvent(new Event('input', { bubbles: true }));
|
||||||
|
|
||||||
|
await vi.runAllTimersAsync();
|
||||||
|
await Promise.resolve();
|
||||||
|
|
||||||
|
const calledUrl = fetchApiMock.mock.calls[0][0];
|
||||||
|
expect(calledUrl).toContain('recursive=true');
|
||||||
|
});
|
||||||
|
|
||||||
|
it('omits folder param when active folder is root and recursion is enabled', async () => {
|
||||||
|
vi.useFakeTimers();
|
||||||
|
|
||||||
|
settingGetMock.mockImplementation((key) => {
|
||||||
|
if (key === 'loramanager.lora_active_filters_autocomplete') {
|
||||||
|
return true;
|
||||||
|
}
|
||||||
|
return undefined;
|
||||||
|
});
|
||||||
|
|
||||||
|
localStorage.setItem('lora_manager_loras_filters', JSON.stringify({
|
||||||
|
baseModel: ['SD 1.5'],
|
||||||
|
tags: { anime: 'include' },
|
||||||
|
}));
|
||||||
|
localStorage.setItem('lora_manager_loras_activeFolder', '');
|
||||||
|
localStorage.removeItem('lora_manager_loras_recursiveSearch');
|
||||||
|
|
||||||
|
fetchApiMock.mockResolvedValue({
|
||||||
|
json: () => Promise.resolve({ success: true, relative_paths: ['models/example.safetensors'] }),
|
||||||
|
});
|
||||||
|
|
||||||
|
caretHelperInstance.getBeforeCursor.mockReturnValue('example');
|
||||||
|
caretHelperInstance.getCursorOffset.mockReturnValue({ left: 15, top: 25 });
|
||||||
|
|
||||||
|
const input = document.createElement('textarea');
|
||||||
|
document.body.append(input);
|
||||||
|
|
||||||
|
const { AutoComplete } = await import(AUTOCOMPLETE_MODULE);
|
||||||
|
new AutoComplete(input, 'loras', { debounceDelay: 0, showPreview: false, minChars: 1 });
|
||||||
|
|
||||||
|
input.value = 'example';
|
||||||
|
input.dispatchEvent(new Event('input', { bubbles: true }));
|
||||||
|
|
||||||
|
await vi.runAllTimersAsync();
|
||||||
|
await Promise.resolve();
|
||||||
|
|
||||||
|
const calledUrl = fetchApiMock.mock.calls[0][0];
|
||||||
|
expect(calledUrl).not.toContain('folder=');
|
||||||
|
expect(calledUrl).toContain('recursive=true');
|
||||||
|
});
|
||||||
|
|
||||||
|
it('sends an empty folder param for root with recursion disabled, mirroring the page list', async () => {
|
||||||
|
vi.useFakeTimers();
|
||||||
|
|
||||||
|
settingGetMock.mockImplementation((key) => {
|
||||||
|
if (key === 'loramanager.lora_active_filters_autocomplete') {
|
||||||
|
return true;
|
||||||
|
}
|
||||||
|
return undefined;
|
||||||
|
});
|
||||||
|
|
||||||
|
localStorage.setItem('lora_manager_loras_filters', JSON.stringify({
|
||||||
|
baseModel: ['SD 1.5'],
|
||||||
|
tags: { anime: 'include' },
|
||||||
|
}));
|
||||||
|
localStorage.setItem('lora_manager_loras_activeFolder', '');
|
||||||
|
localStorage.setItem('lora_manager_loras_recursiveSearch', 'false');
|
||||||
|
|
||||||
|
fetchApiMock.mockResolvedValue({
|
||||||
|
json: () => Promise.resolve({ success: true, relative_paths: ['models/example.safetensors'] }),
|
||||||
|
});
|
||||||
|
|
||||||
|
caretHelperInstance.getBeforeCursor.mockReturnValue('example');
|
||||||
|
caretHelperInstance.getCursorOffset.mockReturnValue({ left: 15, top: 25 });
|
||||||
|
|
||||||
|
const input = document.createElement('textarea');
|
||||||
|
document.body.append(input);
|
||||||
|
|
||||||
|
const { AutoComplete } = await import(AUTOCOMPLETE_MODULE);
|
||||||
|
new AutoComplete(input, 'loras', { debounceDelay: 0, showPreview: false, minChars: 1 });
|
||||||
|
|
||||||
|
input.value = 'example';
|
||||||
|
input.dispatchEvent(new Event('input', { bubbles: true }));
|
||||||
|
|
||||||
|
await vi.runAllTimersAsync();
|
||||||
|
await Promise.resolve();
|
||||||
|
|
||||||
|
const calledUrl = fetchApiMock.mock.calls[0][0];
|
||||||
|
expect(calledUrl).toContain('folder=');
|
||||||
|
expect(calledUrl).toContain('recursive=false');
|
||||||
|
const parsed = new URL(calledUrl, 'https://example.com');
|
||||||
|
expect(parsed.searchParams.get('folder')).toBe('');
|
||||||
|
});
|
||||||
|
|
||||||
|
it('applies the active folder even when no filter-panel filters are set', async () => {
|
||||||
|
// Regression: folder was skipped when lora_manager_loras_filters was
|
||||||
|
// missing because the filters key gate returned early.
|
||||||
|
vi.useFakeTimers();
|
||||||
|
|
||||||
|
settingGetMock.mockImplementation((key) => {
|
||||||
|
if (key === 'loramanager.lora_active_filters_autocomplete') {
|
||||||
|
return true;
|
||||||
|
}
|
||||||
|
return undefined;
|
||||||
|
});
|
||||||
|
|
||||||
|
localStorage.removeItem('lora_manager_loras_filters');
|
||||||
|
localStorage.setItem('lora_manager_loras_activeFolder', 'Flux.1 D/style');
|
||||||
|
|
||||||
|
fetchApiMock.mockResolvedValue({
|
||||||
|
json: () => Promise.resolve({ success: true, relative_paths: ['Flux.1 D/style/3D_Fairytales.safetensors'] }),
|
||||||
|
});
|
||||||
|
|
||||||
|
caretHelperInstance.getBeforeCursor.mockReturnValue('3D');
|
||||||
|
caretHelperInstance.getCursorOffset.mockReturnValue({ left: 15, top: 25 });
|
||||||
|
|
||||||
|
const input = document.createElement('textarea');
|
||||||
|
document.body.append(input);
|
||||||
|
|
||||||
|
const { AutoComplete } = await import(AUTOCOMPLETE_MODULE);
|
||||||
|
new AutoComplete(input, 'loras', { debounceDelay: 0, showPreview: false, minChars: 1 });
|
||||||
|
|
||||||
|
input.value = '3D';
|
||||||
|
input.dispatchEvent(new Event('input', { bubbles: true }));
|
||||||
|
|
||||||
|
await vi.runAllTimersAsync();
|
||||||
|
await Promise.resolve();
|
||||||
|
|
||||||
|
const calledUrl = fetchApiMock.mock.calls[0][0];
|
||||||
|
expect(calledUrl).toContain('folder=Flux.1+D%2Fstyle');
|
||||||
|
expect(calledUrl).toContain('recursive=true');
|
||||||
|
});
|
||||||
});
|
});
|
||||||
|
|||||||
@@ -2155,4 +2155,35 @@ describe('Interaction-level regression coverage', () => {
|
|||||||
excludedItem.dispatchEvent(new Event('click', { bubbles: true }));
|
excludedItem.dispatchEvent(new Event('click', { bubbles: true }));
|
||||||
expect(window.pageControls.enterExcludedView).toHaveBeenCalledTimes(1);
|
expect(window.pageControls.enterExcludedView).toHaveBeenCalledTimes(1);
|
||||||
});
|
});
|
||||||
|
|
||||||
|
it('routes single-model example downloads to missing-only and force paths', async () => {
|
||||||
|
document.body.innerHTML = `
|
||||||
|
<div id="loraContextMenu" class="context-menu">
|
||||||
|
<div class="context-menu-item has-submenu" data-has-submenu="download-examples">
|
||||||
|
<div class="context-submenu">
|
||||||
|
<div class="context-menu-item" data-action="download-examples"></div>
|
||||||
|
<div class="context-menu-item" data-action="download-examples-force"></div>
|
||||||
|
</div>
|
||||||
|
</div>
|
||||||
|
</div>
|
||||||
|
`;
|
||||||
|
|
||||||
|
const { LoraContextMenu } = await import('../../../static/js/components/ContextMenu/LoraContextMenu.js');
|
||||||
|
const contextMenu = new LoraContextMenu();
|
||||||
|
|
||||||
|
const card = document.createElement('div');
|
||||||
|
card.className = 'model-card';
|
||||||
|
card.dataset.filepath = '/models/test.safetensors';
|
||||||
|
card.dataset.sha256 = 'abc123hash';
|
||||||
|
document.body.appendChild(card);
|
||||||
|
|
||||||
|
contextMenu.showMenu(100, 100, card);
|
||||||
|
|
||||||
|
document.querySelector('[data-action="download-examples"]').dispatchEvent(new Event('click', { bubbles: true }));
|
||||||
|
expect(downloadExampleImagesApiMock).toHaveBeenCalledWith(['abc123hash'], null, { force: false });
|
||||||
|
|
||||||
|
contextMenu.showMenu(100, 100, card);
|
||||||
|
document.querySelector('[data-action="download-examples-force"]').dispatchEvent(new Event('click', { bubbles: true }));
|
||||||
|
expect(downloadExampleImagesApiMock).toHaveBeenCalledWith(['abc123hash'], null, { force: true });
|
||||||
|
});
|
||||||
});
|
});
|
||||||
|
|||||||
@@ -0,0 +1,438 @@
|
|||||||
|
import { afterEach, beforeEach, describe, expect, it, vi } from 'vitest';
|
||||||
|
|
||||||
|
const {
|
||||||
|
SUMMARY_MODULE,
|
||||||
|
I18N_HELPERS_MODULE,
|
||||||
|
UI_HELPERS_MODULE,
|
||||||
|
} = vi.hoisted(() => ({
|
||||||
|
SUMMARY_MODULE: new URL('../../../static/js/components/DownloadBatchSummaryModal.js', import.meta.url).pathname,
|
||||||
|
I18N_HELPERS_MODULE: new URL('../../../static/js/utils/i18nHelpers.js', import.meta.url).pathname,
|
||||||
|
UI_HELPERS_MODULE: new URL('../../../static/js/utils/uiHelpers.js', import.meta.url).pathname,
|
||||||
|
}));
|
||||||
|
|
||||||
|
const showToastMock = vi.hoisted(() => vi.fn());
|
||||||
|
const openHuggingFaceMock = vi.hoisted(() => vi.fn());
|
||||||
|
|
||||||
|
vi.mock(I18N_HELPERS_MODULE, () => ({
|
||||||
|
translate: vi.fn((_key, _params, fallback) => fallback ?? ''),
|
||||||
|
}));
|
||||||
|
|
||||||
|
vi.mock(UI_HELPERS_MODULE, () => ({
|
||||||
|
showToast: showToastMock,
|
||||||
|
openHuggingFace: openHuggingFaceMock,
|
||||||
|
}));
|
||||||
|
|
||||||
|
// A realistic failure payload from the backend: a JSON envelope whose `error`
|
||||||
|
// field embeds an HTTP status and a nested JSON body (Civitai Early Access).
|
||||||
|
const REAL_ERROR = '{"success": false, "error": "Failed to resolve authenticated Civitai redirect: status=403 body={\\"error\\":\\"Early Access\\",\\"deadline\\":\\"2026-08-12T08:18:36.063Z\\",\\"message\\":\\"This asset is in Early Access. You can use Buzz access it now!\\"}", "download_id": "1786065633067"}';
|
||||||
|
|
||||||
|
// The human-readable error the component should derive from REAL_ERROR.
|
||||||
|
const FORMATTED_REAL_ERROR = 'HTTP 403 — This asset is in Early Access. You can use Buzz access it now!';
|
||||||
|
|
||||||
|
describe('DownloadBatchSummaryModal', () => {
|
||||||
|
let showDownloadBatchSummary;
|
||||||
|
|
||||||
|
beforeEach(async () => {
|
||||||
|
document.body.innerHTML = '';
|
||||||
|
showToastMock.mockClear();
|
||||||
|
openHuggingFaceMock.mockClear();
|
||||||
|
({ showDownloadBatchSummary } = await import(SUMMARY_MODULE));
|
||||||
|
});
|
||||||
|
|
||||||
|
afterEach(() => {
|
||||||
|
document.body.innerHTML = '';
|
||||||
|
delete navigator.clipboard;
|
||||||
|
delete document.execCommand;
|
||||||
|
vi.restoreAllMocks();
|
||||||
|
vi.useRealTimers();
|
||||||
|
});
|
||||||
|
|
||||||
|
it('renders a warning summary with stat cards and a failure table on partial success', () => {
|
||||||
|
showDownloadBatchSummary({
|
||||||
|
total: 3,
|
||||||
|
completed: 2,
|
||||||
|
failedItems: [
|
||||||
|
{ item: { displayName: 'LoraA' }, error: 'timeout' },
|
||||||
|
{ item: { name: 'LoraB' }, error: '404' },
|
||||||
|
],
|
||||||
|
onRetry: vi.fn(),
|
||||||
|
});
|
||||||
|
|
||||||
|
const modal = document.getElementById('downloadBatchSummaryModal');
|
||||||
|
expect(modal).not.toBeNull();
|
||||||
|
expect(modal.querySelector('.summary-header').classList.contains('warning')).toBe(true);
|
||||||
|
|
||||||
|
// Success / Failed / Total stat cards.
|
||||||
|
const statValues = Array.from(modal.querySelectorAll('.stat-card-value')).map(el => el.textContent);
|
||||||
|
expect(statValues).toEqual(['2', '2', '3']);
|
||||||
|
|
||||||
|
const rows = modal.querySelectorAll('.failure-table tbody tr');
|
||||||
|
expect(rows).toHaveLength(2);
|
||||||
|
expect(rows[0].querySelector('.failure-name').textContent).toBe('LoraA');
|
||||||
|
expect(rows[0].querySelector('.failure-error').textContent).toBe('timeout');
|
||||||
|
expect(rows[1].querySelector('.failure-name').textContent).toBe('LoraB');
|
||||||
|
expect(rows[1].querySelector('.failure-error').textContent).toBe('404');
|
||||||
|
|
||||||
|
expect(modal.querySelector('[data-action="retry-failed"]').textContent).toContain('Retry Failed (2)');
|
||||||
|
expect(modal.querySelector('[data-action="copy-report"]')).not.toBeNull();
|
||||||
|
});
|
||||||
|
|
||||||
|
it('renders an error header when every download failed', () => {
|
||||||
|
showDownloadBatchSummary({
|
||||||
|
total: 2,
|
||||||
|
completed: 0,
|
||||||
|
failedItems: [
|
||||||
|
{ item: { displayName: 'LoraA' }, error: 'timeout' },
|
||||||
|
{ item: { displayName: 'LoraB' }, error: '404' },
|
||||||
|
],
|
||||||
|
onRetry: vi.fn(),
|
||||||
|
});
|
||||||
|
|
||||||
|
const modal = document.getElementById('downloadBatchSummaryModal');
|
||||||
|
expect(modal.querySelector('.summary-header').classList.contains('error')).toBe(true);
|
||||||
|
expect(modal.querySelector('.summary-title').textContent).toBe('Download failed');
|
||||||
|
});
|
||||||
|
|
||||||
|
it('renders a success summary without a failure table or retry button', () => {
|
||||||
|
showDownloadBatchSummary({ total: 2, completed: 2, failedItems: [], onRetry: vi.fn() });
|
||||||
|
|
||||||
|
const modal = document.getElementById('downloadBatchSummaryModal');
|
||||||
|
expect(modal.querySelector('.summary-header').classList.contains('success')).toBe(true);
|
||||||
|
expect(modal.querySelector('.failure-table')).toBeNull();
|
||||||
|
expect(modal.querySelector('[data-action="retry-failed"]')).toBeNull();
|
||||||
|
expect(modal.querySelector('.refresh-success-message')).not.toBeNull();
|
||||||
|
});
|
||||||
|
|
||||||
|
it('escapes HTML in failed item names and errors', () => {
|
||||||
|
showDownloadBatchSummary({
|
||||||
|
total: 1,
|
||||||
|
completed: 0,
|
||||||
|
failedItems: [
|
||||||
|
{ item: { name: '<img src=x onerror=alert(1)>', url: 'https://example.com/xss-model' }, error: '<script>bad()</script>' },
|
||||||
|
],
|
||||||
|
onRetry: vi.fn(),
|
||||||
|
});
|
||||||
|
|
||||||
|
const nameCell = document.querySelector('.failure-name');
|
||||||
|
const errorCell = document.querySelector('.failure-error');
|
||||||
|
|
||||||
|
// The URL resolves, so the name renders inside the failure link; the
|
||||||
|
// escaped entities must render back to the literal payload as text...
|
||||||
|
expect(nameCell.querySelector('a.failure-link')).not.toBeNull();
|
||||||
|
expect(nameCell.textContent).toContain('<img src=x onerror=alert(1)>');
|
||||||
|
expect(errorCell.textContent).toContain('<script>bad()</script>');
|
||||||
|
// ...and never as live DOM nodes.
|
||||||
|
expect(document.querySelector('.failure-table img')).toBeNull();
|
||||||
|
expect(document.querySelector('.failure-table script')).toBeNull();
|
||||||
|
expect(nameCell.innerHTML).toContain('<img');
|
||||||
|
});
|
||||||
|
|
||||||
|
it('removes the modal and invokes onRetry with the original failed items', () => {
|
||||||
|
const onRetry = vi.fn();
|
||||||
|
const failedItems = [{ item: { displayName: 'LoraA' }, error: 'timeout' }];
|
||||||
|
showDownloadBatchSummary({ total: 3, completed: 2, failedItems, onRetry });
|
||||||
|
|
||||||
|
document.querySelector('[data-action="retry-failed"]').click();
|
||||||
|
|
||||||
|
expect(document.getElementById('downloadBatchSummaryModal')).toBeNull();
|
||||||
|
expect(onRetry).toHaveBeenCalledTimes(1);
|
||||||
|
expect(onRetry).toHaveBeenCalledWith(failedItems);
|
||||||
|
// Same object references, not copies.
|
||||||
|
expect(onRetry.mock.calls[0][0][0]).toBe(failedItems[0]);
|
||||||
|
});
|
||||||
|
|
||||||
|
it('closes the modal via the close action without retrying', () => {
|
||||||
|
const onRetry = vi.fn();
|
||||||
|
showDownloadBatchSummary({
|
||||||
|
total: 2,
|
||||||
|
completed: 1,
|
||||||
|
failedItems: [{ item: { name: 'LoraA' }, error: 'timeout' }],
|
||||||
|
onRetry,
|
||||||
|
});
|
||||||
|
|
||||||
|
document.querySelector('.cancel-btn[data-action="close-modal"]').click();
|
||||||
|
|
||||||
|
expect(document.getElementById('downloadBatchSummaryModal')).toBeNull();
|
||||||
|
expect(onRetry).not.toHaveBeenCalled();
|
||||||
|
});
|
||||||
|
|
||||||
|
it('copies a plain-text batch report to the clipboard', async () => {
|
||||||
|
const writeText = vi.fn().mockResolvedValue(undefined);
|
||||||
|
Object.defineProperty(navigator, 'clipboard', { value: { writeText }, configurable: true });
|
||||||
|
|
||||||
|
showDownloadBatchSummary({
|
||||||
|
total: 3,
|
||||||
|
completed: 2,
|
||||||
|
failedItems: [
|
||||||
|
{ item: { displayName: 'LoraA', url: 'https://civitai.red/models/111/lora-a?modelVersionId=222' }, error: 'timeout' },
|
||||||
|
{ item: { name: 'LoraB', url: 'https://example.com/lora-b' }, error: '404' },
|
||||||
|
],
|
||||||
|
onRetry: vi.fn(),
|
||||||
|
});
|
||||||
|
|
||||||
|
document.querySelector('[data-action="copy-report"]').click();
|
||||||
|
|
||||||
|
// writeText is invoked synchronously by the click handler.
|
||||||
|
expect(writeText).toHaveBeenCalledTimes(1);
|
||||||
|
const text = writeText.mock.calls[0][0];
|
||||||
|
expect(text).toContain('Batch Download Report');
|
||||||
|
expect(text).toContain('Total: 3');
|
||||||
|
expect(text).toContain('LoraA — timeout');
|
||||||
|
expect(text).toContain('LoraB — 404');
|
||||||
|
|
||||||
|
// Each failed item with a URL gets an indented URL line right after it.
|
||||||
|
expect(text).toContain(' URL: https://civitai.red/models/111/lora-a?modelVersionId=222');
|
||||||
|
expect(text).toContain(' URL: https://example.com/lora-b');
|
||||||
|
// Exactly the two URLs from the failed items — nothing more, no undefined.
|
||||||
|
expect(text.match(/^\s+URL:/gm)).toHaveLength(2);
|
||||||
|
expect(text).not.toContain('URL: undefined');
|
||||||
|
|
||||||
|
// The toast fires after the mocked clipboard promise settles.
|
||||||
|
await vi.waitFor(() => expect(showToastMock).toHaveBeenCalledTimes(1));
|
||||||
|
expect(showToastMock).toHaveBeenCalledWith('toast.api.copiedToClipboard', {}, 'success');
|
||||||
|
});
|
||||||
|
|
||||||
|
it('omits the URL line for failed items without a resolvable url', async () => {
|
||||||
|
const writeText = vi.fn().mockResolvedValue(undefined);
|
||||||
|
Object.defineProperty(navigator, 'clipboard', { value: { writeText }, configurable: true });
|
||||||
|
|
||||||
|
showDownloadBatchSummary({
|
||||||
|
total: 2,
|
||||||
|
completed: 0,
|
||||||
|
failedItems: [
|
||||||
|
{ item: { name: 'WithUrl', url: 'https://example.com/with-url' }, error: 'boom' },
|
||||||
|
{ item: { name: 'NoUrl' }, error: 'boom' },
|
||||||
|
],
|
||||||
|
onRetry: vi.fn(),
|
||||||
|
});
|
||||||
|
|
||||||
|
document.querySelector('[data-action="copy-report"]').click();
|
||||||
|
|
||||||
|
const text = writeText.mock.calls[0][0];
|
||||||
|
expect(text).toContain(' URL: https://example.com/with-url');
|
||||||
|
// Only the one URL line exists — the URL-less item contributes none.
|
||||||
|
expect(text.match(/^\s+URL:/gm)).toHaveLength(1);
|
||||||
|
expect(text).not.toContain(' URL: undefined');
|
||||||
|
expect(text).not.toContain(' URL: null');
|
||||||
|
|
||||||
|
await vi.waitFor(() => expect(showToastMock).toHaveBeenCalledTimes(1));
|
||||||
|
});
|
||||||
|
|
||||||
|
it('falls back to execCommand when navigator.clipboard is unavailable', async () => {
|
||||||
|
// afterEach deletes navigator.clipboard, but be explicit so this test is
|
||||||
|
// robust even if a previous test failed before its cleanup ran.
|
||||||
|
delete navigator.clipboard;
|
||||||
|
// jsdom does not implement document.execCommand, so install a mock for the
|
||||||
|
// fallback path (removed by the afterEach cleanup above).
|
||||||
|
const execCommandMock = vi.fn(() => true);
|
||||||
|
document.execCommand = execCommandMock;
|
||||||
|
|
||||||
|
showDownloadBatchSummary({
|
||||||
|
total: 3,
|
||||||
|
completed: 2,
|
||||||
|
failedItems: [
|
||||||
|
{ item: { displayName: 'LoraA' }, error: 'timeout' },
|
||||||
|
{ item: { name: 'LoraB' }, error: '404' },
|
||||||
|
],
|
||||||
|
onRetry: vi.fn(),
|
||||||
|
});
|
||||||
|
|
||||||
|
document.querySelector('[data-action="copy-report"]').click();
|
||||||
|
|
||||||
|
// Without the async Clipboard API the fallback must run synchronously.
|
||||||
|
expect(execCommandMock).toHaveBeenCalledWith('copy');
|
||||||
|
|
||||||
|
await Promise.resolve();
|
||||||
|
await Promise.resolve();
|
||||||
|
|
||||||
|
expect(showToastMock).toHaveBeenCalledWith('toast.api.copiedToClipboard', {}, 'success');
|
||||||
|
});
|
||||||
|
|
||||||
|
it('keeps only a single modal instance across repeated calls', () => {
|
||||||
|
showDownloadBatchSummary({
|
||||||
|
total: 2,
|
||||||
|
completed: 1,
|
||||||
|
failedItems: [{ item: { name: 'A' }, error: 'e' }],
|
||||||
|
onRetry: vi.fn(),
|
||||||
|
});
|
||||||
|
showDownloadBatchSummary({ total: 3, completed: 3, failedItems: [], onRetry: vi.fn() });
|
||||||
|
|
||||||
|
expect(document.querySelectorAll('#downloadBatchSummaryModal')).toHaveLength(1);
|
||||||
|
const modal = document.getElementById('downloadBatchSummaryModal');
|
||||||
|
expect(modal.querySelector('.summary-header').classList.contains('success')).toBe(true);
|
||||||
|
});
|
||||||
|
|
||||||
|
it('resolves failure names from entry.name, item fields, URL paths, or Unknown', () => {
|
||||||
|
showDownloadBatchSummary({
|
||||||
|
total: 4,
|
||||||
|
completed: 0,
|
||||||
|
failedItems: [
|
||||||
|
{ name: 'entryName', item: { displayName: 'ItemName' }, error: 'e1' },
|
||||||
|
{ item: { selectedVersion: { name: 'v1.0' } }, error: 'e2' },
|
||||||
|
{ item: { url: 'https://civitai.red/models/837884/midjourney-artful-nsfw?modelVersionId=3153960' }, error: 'e3' },
|
||||||
|
{ item: {}, error: 'e4' },
|
||||||
|
],
|
||||||
|
onRetry: vi.fn(),
|
||||||
|
});
|
||||||
|
|
||||||
|
const names = Array.from(document.querySelectorAll('.failure-name')).map(el => el.textContent);
|
||||||
|
expect(names).toEqual(['entryName', 'v1.0', 'midjourney-artful-nsfw', 'Unknown']);
|
||||||
|
});
|
||||||
|
|
||||||
|
it('formats the real JSON failure payload into a concise HTTP error and truncates long ones', () => {
|
||||||
|
showDownloadBatchSummary({
|
||||||
|
total: 2,
|
||||||
|
completed: 0,
|
||||||
|
failedItems: [
|
||||||
|
{ item: { name: 'EarlyAccess' }, error: REAL_ERROR },
|
||||||
|
{ item: { name: 'LongError' }, error: 'x'.repeat(300) },
|
||||||
|
],
|
||||||
|
onRetry: vi.fn(),
|
||||||
|
});
|
||||||
|
|
||||||
|
const errorCells = document.querySelectorAll('.failure-error');
|
||||||
|
expect(errorCells[0].textContent).toBe(FORMATTED_REAL_ERROR);
|
||||||
|
expect(errorCells[1].textContent).toBe('x'.repeat(220) + '…');
|
||||||
|
});
|
||||||
|
|
||||||
|
it('keeps the raw error string in the error cell title for debugging', () => {
|
||||||
|
showDownloadBatchSummary({
|
||||||
|
total: 1,
|
||||||
|
completed: 0,
|
||||||
|
failedItems: [{ item: { name: 'EarlyAccess' }, error: REAL_ERROR }],
|
||||||
|
onRetry: vi.fn(),
|
||||||
|
});
|
||||||
|
|
||||||
|
const errorCell = document.querySelector('.failure-error');
|
||||||
|
expect(errorCell.getAttribute('title')).toBe(REAL_ERROR);
|
||||||
|
expect(errorCell.getAttribute('title')).not.toBe(FORMATTED_REAL_ERROR);
|
||||||
|
});
|
||||||
|
|
||||||
|
it('opens the original item url in a new tab when a failure link is clicked', () => {
|
||||||
|
showDownloadBatchSummary({
|
||||||
|
total: 1,
|
||||||
|
completed: 0,
|
||||||
|
failedItems: [{
|
||||||
|
item: {
|
||||||
|
url: 'https://civitai.red/models/837884/midjourney-artful-nsfw?modelVersionId=3153960',
|
||||||
|
modelId: '837884',
|
||||||
|
selectedVersion: { id: '3153960' },
|
||||||
|
},
|
||||||
|
error: 'rate limited',
|
||||||
|
}],
|
||||||
|
onRetry: vi.fn(),
|
||||||
|
});
|
||||||
|
|
||||||
|
document.querySelector('.failure-link').click();
|
||||||
|
|
||||||
|
expect(openHuggingFaceMock).toHaveBeenCalledTimes(1);
|
||||||
|
expect(openHuggingFaceMock).toHaveBeenCalledWith('https://civitai.red/models/837884/midjourney-artful-nsfw?modelVersionId=3153960');
|
||||||
|
// The modal stays open so the user can keep inspecting the failures.
|
||||||
|
expect(document.getElementById('downloadBatchSummaryModal')).not.toBeNull();
|
||||||
|
});
|
||||||
|
|
||||||
|
it('opens the item url directly when selectedVersion is absent', () => {
|
||||||
|
showDownloadBatchSummary({
|
||||||
|
total: 1,
|
||||||
|
completed: 0,
|
||||||
|
failedItems: [{
|
||||||
|
item: {
|
||||||
|
modelId: '837884',
|
||||||
|
modelVersionId: '3153960',
|
||||||
|
url: 'https://civitai.red/models/837884/midjourney-artful-nsfw',
|
||||||
|
},
|
||||||
|
error: 'rate limited',
|
||||||
|
}],
|
||||||
|
onRetry: vi.fn(),
|
||||||
|
});
|
||||||
|
|
||||||
|
document.querySelector('.failure-link').click();
|
||||||
|
|
||||||
|
expect(openHuggingFaceMock).toHaveBeenCalledTimes(1);
|
||||||
|
expect(openHuggingFaceMock).toHaveBeenCalledWith('https://civitai.red/models/837884/midjourney-artful-nsfw');
|
||||||
|
expect(document.getElementById('downloadBatchSummaryModal')).not.toBeNull();
|
||||||
|
});
|
||||||
|
|
||||||
|
it('opens the original huggingface url directly when a huggingface failure link is clicked', () => {
|
||||||
|
showDownloadBatchSummary({
|
||||||
|
total: 1,
|
||||||
|
completed: 0,
|
||||||
|
failedItems: [{
|
||||||
|
item: {
|
||||||
|
url: 'https://huggingface.co/user/repo',
|
||||||
|
source: 'huggingface',
|
||||||
|
repo: 'user/repo',
|
||||||
|
filename: 'model.safetensors',
|
||||||
|
revision: 'main',
|
||||||
|
},
|
||||||
|
error: 'download failed',
|
||||||
|
}],
|
||||||
|
onRetry: vi.fn(),
|
||||||
|
});
|
||||||
|
|
||||||
|
document.querySelector('.failure-link').click();
|
||||||
|
|
||||||
|
expect(openHuggingFaceMock).toHaveBeenCalledTimes(1);
|
||||||
|
expect(openHuggingFaceMock).toHaveBeenCalledWith('https://huggingface.co/user/repo');
|
||||||
|
expect(document.getElementById('downloadBatchSummaryModal')).not.toBeNull();
|
||||||
|
});
|
||||||
|
|
||||||
|
it('opens an arbitrary URL via openHuggingFace for fallback items', () => {
|
||||||
|
showDownloadBatchSummary({
|
||||||
|
total: 1,
|
||||||
|
completed: 0,
|
||||||
|
failedItems: [{ item: { url: 'https://example.com/model' }, error: 'boom' }],
|
||||||
|
onRetry: vi.fn(),
|
||||||
|
});
|
||||||
|
|
||||||
|
document.querySelector('.failure-link').click();
|
||||||
|
|
||||||
|
expect(openHuggingFaceMock).toHaveBeenCalledTimes(1);
|
||||||
|
expect(openHuggingFaceMock).toHaveBeenCalledWith('https://example.com/model');
|
||||||
|
});
|
||||||
|
|
||||||
|
it('renders the failure name as plain text when no URL can be resolved', () => {
|
||||||
|
showDownloadBatchSummary({
|
||||||
|
total: 1,
|
||||||
|
completed: 0,
|
||||||
|
failedItems: [{ item: { modelId: null }, error: 'boom' }],
|
||||||
|
onRetry: vi.fn(),
|
||||||
|
});
|
||||||
|
|
||||||
|
expect(document.querySelector('a.failure-link')).toBeNull();
|
||||||
|
expect(document.querySelector('.failure-name').textContent).toBe('Unknown');
|
||||||
|
|
||||||
|
// Without a link there is nothing to open: clicking the cell is inert.
|
||||||
|
document.querySelector('.failure-name').click();
|
||||||
|
expect(openHuggingFaceMock).not.toHaveBeenCalled();
|
||||||
|
});
|
||||||
|
|
||||||
|
it('copies formatted errors (not raw JSON) into the report text', async () => {
|
||||||
|
const writeText = vi.fn().mockResolvedValue(undefined);
|
||||||
|
Object.defineProperty(navigator, 'clipboard', { value: { writeText }, configurable: true });
|
||||||
|
|
||||||
|
showDownloadBatchSummary({
|
||||||
|
total: 1,
|
||||||
|
completed: 0,
|
||||||
|
failedItems: [{
|
||||||
|
item: {
|
||||||
|
name: 'EarlyAccess',
|
||||||
|
url: 'https://civitai.red/models/123/early-access?modelVersionId=456',
|
||||||
|
},
|
||||||
|
error: REAL_ERROR,
|
||||||
|
}],
|
||||||
|
onRetry: vi.fn(),
|
||||||
|
});
|
||||||
|
|
||||||
|
document.querySelector('[data-action="copy-report"]').click();
|
||||||
|
|
||||||
|
expect(writeText).toHaveBeenCalledTimes(1);
|
||||||
|
const text = writeText.mock.calls[0][0];
|
||||||
|
expect(text).toContain(FORMATTED_REAL_ERROR);
|
||||||
|
expect(text).toContain(' URL: https://civitai.red/models/123/early-access?modelVersionId=456');
|
||||||
|
expect(text).not.toContain('download_id');
|
||||||
|
expect(text).not.toContain('Failed to resolve authenticated Civitai redirect');
|
||||||
|
|
||||||
|
await vi.waitFor(() => expect(showToastMock).toHaveBeenCalledTimes(1));
|
||||||
|
});
|
||||||
|
});
|
||||||
@@ -0,0 +1,87 @@
|
|||||||
|
import { describe, it, beforeEach, afterEach, expect, vi } from 'vitest';
|
||||||
|
|
||||||
|
const showToastMock = vi.fn();
|
||||||
|
const recreateVirtualScrollMock = vi.fn();
|
||||||
|
|
||||||
|
vi.mock('../../../static/js/utils/uiHelpers.js', () => ({
|
||||||
|
showToast: showToastMock,
|
||||||
|
}));
|
||||||
|
|
||||||
|
vi.mock('../../../static/js/components/RecipeCard.js', () => ({
|
||||||
|
RecipeCard: class {},
|
||||||
|
}));
|
||||||
|
|
||||||
|
vi.mock('../../../static/js/utils/infiniteScroll.js', () => ({
|
||||||
|
recreateVirtualScroll: recreateVirtualScrollMock,
|
||||||
|
}));
|
||||||
|
|
||||||
|
const { DuplicatesManager } = await import('../../../static/js/components/DuplicatesManager.js');
|
||||||
|
const { state, getCurrentPageState, setCurrentPageType } = await import('../../../static/js/state/index.js');
|
||||||
|
|
||||||
|
function setupDom() {
|
||||||
|
document.body.innerHTML = `
|
||||||
|
<div id="duplicatesBanner" style="display: block;"></div>
|
||||||
|
<div id="recipeGrid"><div class="model-card">stale</div></div>
|
||||||
|
`;
|
||||||
|
document.body.classList.add('duplicate-mode');
|
||||||
|
}
|
||||||
|
|
||||||
|
describe('DuplicatesManager exitDuplicateMode', () => {
|
||||||
|
beforeEach(() => {
|
||||||
|
vi.clearAllMocks();
|
||||||
|
setCurrentPageType('recipes');
|
||||||
|
setupDom();
|
||||||
|
state.pendingLayoutRecreate = false;
|
||||||
|
state.virtualScroller = { enable: vi.fn(), disable: vi.fn() };
|
||||||
|
});
|
||||||
|
|
||||||
|
afterEach(() => {
|
||||||
|
state.pendingLayoutRecreate = false;
|
||||||
|
state.virtualScroller = null;
|
||||||
|
});
|
||||||
|
|
||||||
|
it('skips enable() on the old scroller when a layout recreate was deferred', async () => {
|
||||||
|
state.pendingLayoutRecreate = true;
|
||||||
|
|
||||||
|
const manager = new DuplicatesManager({});
|
||||||
|
manager.inDuplicateMode = true;
|
||||||
|
await manager.exitDuplicateMode();
|
||||||
|
|
||||||
|
expect(state.virtualScroller.enable).not.toHaveBeenCalled();
|
||||||
|
expect(recreateVirtualScrollMock).toHaveBeenCalledWith('recipes');
|
||||||
|
expect(state.pendingLayoutRecreate).toBe(false);
|
||||||
|
});
|
||||||
|
|
||||||
|
it('re-enables the existing scroller when no layout recreate is pending', async () => {
|
||||||
|
const manager = new DuplicatesManager({});
|
||||||
|
manager.inDuplicateMode = true;
|
||||||
|
await manager.exitDuplicateMode();
|
||||||
|
|
||||||
|
expect(state.virtualScroller.enable).toHaveBeenCalledTimes(1);
|
||||||
|
expect(recreateVirtualScrollMock).not.toHaveBeenCalled();
|
||||||
|
});
|
||||||
|
|
||||||
|
it('tolerates a missing scroller on the plain re-enable path', async () => {
|
||||||
|
state.virtualScroller = null;
|
||||||
|
|
||||||
|
const manager = new DuplicatesManager({});
|
||||||
|
manager.inDuplicateMode = true;
|
||||||
|
await expect(manager.exitDuplicateMode()).resolves.toBeUndefined();
|
||||||
|
|
||||||
|
expect(recreateVirtualScrollMock).not.toHaveBeenCalled();
|
||||||
|
});
|
||||||
|
|
||||||
|
it('clears duplicates-mode state and the grid regardless of path', async () => {
|
||||||
|
state.pendingLayoutRecreate = true;
|
||||||
|
|
||||||
|
const manager = new DuplicatesManager({});
|
||||||
|
manager.inDuplicateMode = true;
|
||||||
|
await manager.exitDuplicateMode();
|
||||||
|
|
||||||
|
expect(manager.inDuplicateMode).toBe(false);
|
||||||
|
expect(getCurrentPageState().duplicatesMode).toBe(false);
|
||||||
|
expect(document.body.classList.contains('duplicate-mode')).toBe(false);
|
||||||
|
expect(document.getElementById('recipeGrid').innerHTML).toBe('');
|
||||||
|
expect(document.getElementById('duplicatesBanner').style.display).toBe('none');
|
||||||
|
});
|
||||||
|
});
|
||||||
@@ -0,0 +1,221 @@
|
|||||||
|
import { describe, it, beforeEach, afterEach, expect, vi } from 'vitest';
|
||||||
|
|
||||||
|
const resetAndReloadMock = vi.fn();
|
||||||
|
const getModelApiClientMock = vi.fn();
|
||||||
|
|
||||||
|
vi.mock('../../../static/js/api/modelApiFactory.js', () => ({
|
||||||
|
getModelApiClient: getModelApiClientMock,
|
||||||
|
resetAndReload: resetAndReloadMock,
|
||||||
|
}));
|
||||||
|
|
||||||
|
vi.mock('../../../static/js/utils/uiHelpers.js', () => ({
|
||||||
|
showToast: vi.fn(),
|
||||||
|
openCivitaiByMetadata: vi.fn(),
|
||||||
|
updatePanelPositions: vi.fn(),
|
||||||
|
}));
|
||||||
|
|
||||||
|
vi.mock('../../../static/js/managers/DownloadManager.js', () => ({
|
||||||
|
downloadManager: { showDownloadModal: vi.fn() },
|
||||||
|
}));
|
||||||
|
|
||||||
|
vi.mock('../../../static/js/components/SidebarManager.js', () => ({
|
||||||
|
sidebarManager: {
|
||||||
|
setHostPageControls: vi.fn(),
|
||||||
|
initialize: vi.fn(async () => {}),
|
||||||
|
refresh: vi.fn(async () => {}),
|
||||||
|
cleanup: vi.fn(),
|
||||||
|
isInitialized: false,
|
||||||
|
},
|
||||||
|
}));
|
||||||
|
|
||||||
|
vi.mock('../../../static/js/components/alphabet/index.js', () => ({
|
||||||
|
createAlphabetBar: vi.fn(() => ({ destroy: vi.fn() })),
|
||||||
|
}));
|
||||||
|
|
||||||
|
vi.mock('../../../static/js/utils/updateCheckHelpers.js', () => ({
|
||||||
|
performModelUpdateCheck: vi.fn(async () => ({ status: 'success', displayName: 'LoRA', records: [] })),
|
||||||
|
}));
|
||||||
|
|
||||||
|
beforeEach(() => {
|
||||||
|
vi.resetModules();
|
||||||
|
vi.clearAllMocks();
|
||||||
|
localStorage.clear();
|
||||||
|
sessionStorage.clear();
|
||||||
|
|
||||||
|
resetAndReloadMock.mockResolvedValue(undefined);
|
||||||
|
getModelApiClientMock.mockReturnValue({});
|
||||||
|
|
||||||
|
global.fetch = vi.fn().mockResolvedValue({
|
||||||
|
ok: true,
|
||||||
|
json: async () => ({ success: true, base_models: [] }),
|
||||||
|
});
|
||||||
|
});
|
||||||
|
|
||||||
|
afterEach(() => {
|
||||||
|
delete window.bulkManager;
|
||||||
|
delete window.modelDuplicatesManager;
|
||||||
|
delete global.fetch;
|
||||||
|
});
|
||||||
|
|
||||||
|
function renderControlsDom(pageKey) {
|
||||||
|
document.body.dataset.page = pageKey;
|
||||||
|
document.body.innerHTML = `
|
||||||
|
<div class="controls">
|
||||||
|
<div id="excludedViewBanner" class="excluded-view-banner hidden">
|
||||||
|
<button id="excludedViewBackBtn">Back</button>
|
||||||
|
</div>
|
||||||
|
<div class="actions">
|
||||||
|
<div class="action-buttons">
|
||||||
|
<div class="control-group">
|
||||||
|
<select id="sortSelect">
|
||||||
|
<option value="name:asc">Name Asc</option>
|
||||||
|
<option value="name:desc">Name Desc</option>
|
||||||
|
<option value="random">Randomize (shuffle)</option>
|
||||||
|
</select>
|
||||||
|
</div>
|
||||||
|
<div class="control-group dropdown-group">
|
||||||
|
<button data-action="refresh" class="dropdown-main"></button>
|
||||||
|
<button class="dropdown-toggle"></button>
|
||||||
|
<div class="dropdown-menu">
|
||||||
|
<div class="dropdown-item" data-action="full-rebuild"></div>
|
||||||
|
</div>
|
||||||
|
</div>
|
||||||
|
<div class="control-group">
|
||||||
|
<button data-action="fetch"></button>
|
||||||
|
</div>
|
||||||
|
<div class="control-group">
|
||||||
|
<button data-action="download"></button>
|
||||||
|
</div>
|
||||||
|
<div class="control-group">
|
||||||
|
<button data-action="bulk"></button>
|
||||||
|
</div>
|
||||||
|
<div class="control-group">
|
||||||
|
<button data-action="find-duplicates"></button>
|
||||||
|
</div>
|
||||||
|
<div class="control-group">
|
||||||
|
<button id="favoriteFilterBtn" class="favorite-filter"></button>
|
||||||
|
</div>
|
||||||
|
<div class="control-group dropdown-group update-filter-group">
|
||||||
|
<button id="updateFilterBtn" class="dropdown-main update-filter" aria-busy="false">
|
||||||
|
<span>Updates</span>
|
||||||
|
</button>
|
||||||
|
<button id="updateFilterMenuToggle" class="dropdown-toggle"></button>
|
||||||
|
<div class="dropdown-menu">
|
||||||
|
<div id="checkUpdatesMenuItem" class="dropdown-item" data-action="check-updates">
|
||||||
|
<span>Check updates</span>
|
||||||
|
</div>
|
||||||
|
</div>
|
||||||
|
</div>
|
||||||
|
</div>
|
||||||
|
</div>
|
||||||
|
</div>
|
||||||
|
<div id="customFilterIndicator" class="control-group hidden">
|
||||||
|
<div class="filter-active">
|
||||||
|
<span class="customFilterText" title=""></span>
|
||||||
|
<i class="fas fa-times-circle clear-filter"></i>
|
||||||
|
</div>
|
||||||
|
</div>
|
||||||
|
<div id="breadcrumbContainer"></div>
|
||||||
|
<div id="duplicatesBanner" style="display: none;"></div>
|
||||||
|
<div class="alphabet-bar-container"></div>
|
||||||
|
`;
|
||||||
|
}
|
||||||
|
|
||||||
|
async function createControls() {
|
||||||
|
const stateModule = await import('../../../static/js/state/index.js');
|
||||||
|
stateModule.initPageState('loras');
|
||||||
|
const { LorasControls } = await import('../../../static/js/components/controls/LorasControls.js');
|
||||||
|
return { stateModule, controls: new LorasControls() };
|
||||||
|
}
|
||||||
|
|
||||||
|
describe('Random sort option', () => {
|
||||||
|
it('generates a seeded sort value when Random is picked', async () => {
|
||||||
|
renderControlsDom('loras');
|
||||||
|
const { controls } = await createControls();
|
||||||
|
const sortSelect = document.getElementById('sortSelect');
|
||||||
|
const randomOpt = sortSelect.querySelector('option[value="random"]');
|
||||||
|
|
||||||
|
sortSelect.value = 'random';
|
||||||
|
sortSelect.dispatchEvent(new Event('change', { bubbles: true }));
|
||||||
|
await Promise.resolve();
|
||||||
|
|
||||||
|
expect(controls.pageState.sortBy).toMatch(/^random:[a-z0-9]+$/);
|
||||||
|
expect(localStorage.getItem('lora_manager_loras_sort')).toBe(controls.pageState.sortBy);
|
||||||
|
expect(randomOpt.value).toBe(controls.pageState.sortBy);
|
||||||
|
expect(sortSelect.value).toBe(controls.pageState.sortBy);
|
||||||
|
expect(resetAndReloadMock).toHaveBeenCalled();
|
||||||
|
});
|
||||||
|
|
||||||
|
it('reshuffles with a fresh seed every time Random is picked again', async () => {
|
||||||
|
renderControlsDom('loras');
|
||||||
|
const { controls } = await createControls();
|
||||||
|
const sortSelect = document.getElementById('sortSelect');
|
||||||
|
const randomOpt = sortSelect.querySelector('option[value="random"]');
|
||||||
|
|
||||||
|
// First pick
|
||||||
|
sortSelect.value = 'random';
|
||||||
|
sortSelect.dispatchEvent(new Event('change', { bubbles: true }));
|
||||||
|
await Promise.resolve();
|
||||||
|
const firstSeed = controls.pageState.sortBy;
|
||||||
|
|
||||||
|
// Second pick: the option now carries the seeded value, like a menu click
|
||||||
|
sortSelect.value = randomOpt.value;
|
||||||
|
sortSelect.dispatchEvent(new Event('change', { bubbles: true }));
|
||||||
|
await Promise.resolve();
|
||||||
|
|
||||||
|
expect(controls.pageState.sortBy).toMatch(/^random:[a-z0-9]+$/);
|
||||||
|
expect(controls.pageState.sortBy).not.toBe(firstSeed);
|
||||||
|
});
|
||||||
|
|
||||||
|
it('restores a persisted seeded random sort on load', async () => {
|
||||||
|
renderControlsDom('loras');
|
||||||
|
const savedSort = 'random:persistedseed';
|
||||||
|
localStorage.setItem('lora_manager_loras_sort', savedSort);
|
||||||
|
|
||||||
|
const { controls } = await createControls();
|
||||||
|
const sortSelect = document.getElementById('sortSelect');
|
||||||
|
|
||||||
|
expect(controls.pageState.sortBy).toBe(savedSort);
|
||||||
|
expect(sortSelect.value).toBe(savedSort);
|
||||||
|
expect(sortSelect.querySelector('option[value="random:persistedseed"]')).not.toBeNull();
|
||||||
|
});
|
||||||
|
|
||||||
|
it('applies a non-random sort back to the plain random option', async () => {
|
||||||
|
renderControlsDom('loras');
|
||||||
|
const { controls } = await createControls();
|
||||||
|
const sortSelect = document.getElementById('sortSelect');
|
||||||
|
const randomOpt = sortSelect.querySelector('option[value="random"]');
|
||||||
|
|
||||||
|
// Seed a random sort, then switch to a normal sort
|
||||||
|
sortSelect.value = 'random';
|
||||||
|
sortSelect.dispatchEvent(new Event('change', { bubbles: true }));
|
||||||
|
await Promise.resolve();
|
||||||
|
controls.applySortToSelect('name:desc');
|
||||||
|
|
||||||
|
expect(sortSelect.value).toBe('name:desc');
|
||||||
|
expect(randomOpt.value).toBe('random');
|
||||||
|
});
|
||||||
|
|
||||||
|
it('resets the seeded option when switching away from Random via the dropdown change handler', async () => {
|
||||||
|
renderControlsDom('loras');
|
||||||
|
const { controls } = await createControls();
|
||||||
|
const sortSelect = document.getElementById('sortSelect');
|
||||||
|
const randomOpt = sortSelect.querySelector('option[value="random"]');
|
||||||
|
|
||||||
|
// Pick Random: the option is now seeded
|
||||||
|
sortSelect.value = 'random';
|
||||||
|
sortSelect.dispatchEvent(new Event('change', { bubbles: true }));
|
||||||
|
await Promise.resolve();
|
||||||
|
expect(randomOpt.value).toMatch(/^random:[a-z0-9]+$/);
|
||||||
|
|
||||||
|
// Switch to a non-random sort through the change handler (as a menu
|
||||||
|
// click does); the option must go back to the plain "random" value
|
||||||
|
sortSelect.value = 'name:desc';
|
||||||
|
sortSelect.dispatchEvent(new Event('change', { bubbles: true }));
|
||||||
|
await Promise.resolve();
|
||||||
|
|
||||||
|
expect(controls.pageState.sortBy).toBe('name:desc');
|
||||||
|
expect(sortSelect.value).toBe('name:desc');
|
||||||
|
expect(randomOpt.value).toBe('random');
|
||||||
|
});
|
||||||
|
});
|
||||||
@@ -0,0 +1,68 @@
|
|||||||
|
import { describe, it, beforeEach, expect } from 'vitest';
|
||||||
|
import { initSortDropdown } from '../../../static/js/components/controls/SortDropdown.js';
|
||||||
|
|
||||||
|
function renderSortDropdownDom() {
|
||||||
|
document.body.innerHTML = `
|
||||||
|
<div class="sort-dropdown-group">
|
||||||
|
<select id="sortSelect">
|
||||||
|
<option value="name:asc">Name Asc</option>
|
||||||
|
<option value="name:desc">Name Desc</option>
|
||||||
|
<option value="random" selected>Randomize (shuffle)</option>
|
||||||
|
</select>
|
||||||
|
<button class="sort-trigger" type="button">
|
||||||
|
<span class="sort-trigger__label"></span>
|
||||||
|
</button>
|
||||||
|
<div class="sort-dropdown-menu"></div>
|
||||||
|
</div>
|
||||||
|
`;
|
||||||
|
return {
|
||||||
|
select: document.getElementById('sortSelect'),
|
||||||
|
menu: document.querySelector('.sort-dropdown-menu'),
|
||||||
|
label: document.querySelector('.sort-trigger__label'),
|
||||||
|
};
|
||||||
|
}
|
||||||
|
|
||||||
|
describe('SortDropdown menu sync', () => {
|
||||||
|
let select;
|
||||||
|
let menu;
|
||||||
|
let label;
|
||||||
|
|
||||||
|
beforeEach(() => {
|
||||||
|
({ select, menu, label } = renderSortDropdownDom());
|
||||||
|
initSortDropdown(select);
|
||||||
|
});
|
||||||
|
|
||||||
|
it('rebuilds the menu and highlights the selected item when an option value attribute changes', async () => {
|
||||||
|
// The seeded Random option gets a new value each time it is picked.
|
||||||
|
// The select's value getter follows the selected option's new value.
|
||||||
|
const randomOpt = select.querySelector('option[value="random"]');
|
||||||
|
randomOpt.value = 'random:abc123';
|
||||||
|
await Promise.resolve();
|
||||||
|
|
||||||
|
const items = [...menu.querySelectorAll('.sort-option')];
|
||||||
|
expect(items.map((el) => el.dataset.value)).toContain('random:abc123');
|
||||||
|
const seededItem = items.find((el) => el.dataset.value === 'random:abc123');
|
||||||
|
expect(seededItem.classList.contains('is-selected')).toBe(true);
|
||||||
|
expect(label.textContent).toBe('Randomize (shuffle)');
|
||||||
|
});
|
||||||
|
|
||||||
|
it('drops the stale seeded item and re-selects the plain random item when the option is reset', async () => {
|
||||||
|
const randomOpt = select.querySelector('option[value="random"]');
|
||||||
|
randomOpt.value = 'random:abc123';
|
||||||
|
await Promise.resolve();
|
||||||
|
|
||||||
|
// The rebuild must have happened: the seeded item is in the menu
|
||||||
|
const seededItems = [...menu.querySelectorAll('.sort-option')]
|
||||||
|
.filter((el) => el.dataset.value === 'random:abc123');
|
||||||
|
expect(seededItems).toHaveLength(1);
|
||||||
|
|
||||||
|
// PageControls resets the option to "random" when switching away
|
||||||
|
randomOpt.value = 'random';
|
||||||
|
await Promise.resolve();
|
||||||
|
|
||||||
|
const items = [...menu.querySelectorAll('.sort-option')];
|
||||||
|
expect(items.map((el) => el.dataset.value)).not.toContain('random:abc123');
|
||||||
|
const randomItem = items.find((el) => el.dataset.value === 'random');
|
||||||
|
expect(randomItem.classList.contains('is-selected')).toBe(true);
|
||||||
|
});
|
||||||
|
});
|
||||||
@@ -92,6 +92,7 @@ vi.mock('../../../static/js/utils/eventManagementInit.js', () => ({
|
|||||||
|
|
||||||
vi.mock('../../../static/js/utils/infiniteScroll.js', () => ({
|
vi.mock('../../../static/js/utils/infiniteScroll.js', () => ({
|
||||||
initializeInfiniteScroll: vi.fn(),
|
initializeInfiniteScroll: vi.fn(),
|
||||||
|
recreateVirtualScroll: vi.fn(),
|
||||||
}));
|
}));
|
||||||
|
|
||||||
vi.mock('../../../static/js/components/ContextMenu/index.js', () => ({
|
vi.mock('../../../static/js/components/ContextMenu/index.js', () => ({
|
||||||
|
|||||||
@@ -0,0 +1,195 @@
|
|||||||
|
import { beforeEach, describe, expect, it, vi } from "vitest";
|
||||||
|
|
||||||
|
const { APP_MODULE, EXTENSION_MODULE, appMock, registeredExtensions } =
|
||||||
|
vi.hoisted(() => {
|
||||||
|
const registeredExtensions = [];
|
||||||
|
const appMock = {
|
||||||
|
configuringGraph: false,
|
||||||
|
registerExtension: (ext) => registeredExtensions.push(ext),
|
||||||
|
};
|
||||||
|
return {
|
||||||
|
APP_MODULE: new URL("../../../scripts/app.js", import.meta.url).pathname,
|
||||||
|
EXTENSION_MODULE: new URL(
|
||||||
|
"../../../web/comfyui/lora_stack_dynamic_inputs.js",
|
||||||
|
import.meta.url
|
||||||
|
).pathname,
|
||||||
|
appMock,
|
||||||
|
registeredExtensions,
|
||||||
|
};
|
||||||
|
});
|
||||||
|
|
||||||
|
vi.mock(APP_MODULE, () => ({
|
||||||
|
app: appMock,
|
||||||
|
}));
|
||||||
|
|
||||||
|
describe("Lora Stack Combiner dynamic inputs", () => {
|
||||||
|
let extension;
|
||||||
|
|
||||||
|
beforeEach(async () => {
|
||||||
|
vi.resetModules();
|
||||||
|
registeredExtensions.length = 0;
|
||||||
|
appMock.configuringGraph = false;
|
||||||
|
await import(EXTENSION_MODULE);
|
||||||
|
extension = registeredExtensions.find(
|
||||||
|
(ext) => ext.name === "Comfy.LoraManager.LoraStackCombiner"
|
||||||
|
);
|
||||||
|
expect(extension).toBeDefined();
|
||||||
|
});
|
||||||
|
|
||||||
|
function createNodeType() {
|
||||||
|
const nodeType = { prototype: {} };
|
||||||
|
extension.beforeRegisterNodeDef(
|
||||||
|
nodeType,
|
||||||
|
{ name: "Lora Stack Combiner (LoraManager)" },
|
||||||
|
appMock
|
||||||
|
);
|
||||||
|
return nodeType;
|
||||||
|
}
|
||||||
|
|
||||||
|
function createNode(inputs = []) {
|
||||||
|
const node = {
|
||||||
|
comfyClass: "Lora Stack Combiner (LoraManager)",
|
||||||
|
inputs: inputs.map((name) => ({ name, type: "LORA_STACK" })),
|
||||||
|
addInput: vi.fn(function (name, type, opts) {
|
||||||
|
this.inputs.push({ name, type, ...opts });
|
||||||
|
}),
|
||||||
|
removeInput: vi.fn(function (index) {
|
||||||
|
this.inputs.splice(index, 1);
|
||||||
|
}),
|
||||||
|
};
|
||||||
|
return node;
|
||||||
|
}
|
||||||
|
|
||||||
|
function makeLinkInfo() {
|
||||||
|
return { id: 999, origin_id: 1, target_id: 2 };
|
||||||
|
}
|
||||||
|
|
||||||
|
it("adds a third input when the last slot gets connected", () => {
|
||||||
|
const nodeType = createNodeType();
|
||||||
|
const node = createNode(["lora_stack1", "lora_stack2"]);
|
||||||
|
node.onConnectionsChange = nodeType.prototype.onConnectionsChange;
|
||||||
|
|
||||||
|
node.onConnectionsChange(1, 1, true, makeLinkInfo());
|
||||||
|
|
||||||
|
expect(node.inputs.map((input) => input.name)).toEqual([
|
||||||
|
"lora_stack1",
|
||||||
|
"lora_stack2",
|
||||||
|
"lora_stack3",
|
||||||
|
]);
|
||||||
|
});
|
||||||
|
|
||||||
|
it("does not add an input when a non-last slot gets connected", () => {
|
||||||
|
const nodeType = createNodeType();
|
||||||
|
const node = createNode(["lora_stack1", "lora_stack2", "lora_stack3"]);
|
||||||
|
node.onConnectionsChange = nodeType.prototype.onConnectionsChange;
|
||||||
|
|
||||||
|
node.onConnectionsChange(1, 0, true, makeLinkInfo());
|
||||||
|
|
||||||
|
expect(node.inputs.map((input) => input.name)).toEqual([
|
||||||
|
"lora_stack1",
|
||||||
|
"lora_stack2",
|
||||||
|
"lora_stack3",
|
||||||
|
]);
|
||||||
|
});
|
||||||
|
|
||||||
|
it("removes a disconnected middle slot and renumbers", () => {
|
||||||
|
// Simulates a real LiteGraph disconnect event: it fires only for slots that
|
||||||
|
// had a link, and input.link has already been cleared before the event fires.
|
||||||
|
const nodeType = createNodeType();
|
||||||
|
const node = createNode(["lora_stack1", "lora_stack2", "lora_stack3"]);
|
||||||
|
node.inputs[0].link = 11;
|
||||||
|
node.inputs[1].link = null; // slot 2 was just disconnected
|
||||||
|
node.inputs[2].link = 13;
|
||||||
|
node.onConnectionsChange = nodeType.prototype.onConnectionsChange;
|
||||||
|
|
||||||
|
node.onConnectionsChange(1, 1, false, makeLinkInfo());
|
||||||
|
|
||||||
|
expect(node.inputs.map((input) => input.name)).toEqual([
|
||||||
|
"lora_stack1",
|
||||||
|
"lora_stack2",
|
||||||
|
]);
|
||||||
|
});
|
||||||
|
|
||||||
|
it("keeps the last slot when it is disconnected", () => {
|
||||||
|
const nodeType = createNodeType();
|
||||||
|
const node = createNode(["lora_stack1", "lora_stack2", "lora_stack3"]);
|
||||||
|
node.inputs[0].link = 11;
|
||||||
|
node.inputs[1].link = 12;
|
||||||
|
node.inputs[2].link = null; // last slot was just disconnected
|
||||||
|
node.onConnectionsChange = nodeType.prototype.onConnectionsChange;
|
||||||
|
|
||||||
|
node.onConnectionsChange(1, 2, false, makeLinkInfo());
|
||||||
|
|
||||||
|
expect(node.inputs.map((input) => input.name)).toEqual([
|
||||||
|
"lora_stack1",
|
||||||
|
"lora_stack2",
|
||||||
|
"lora_stack3",
|
||||||
|
]);
|
||||||
|
expect(node.removeInput).not.toHaveBeenCalled();
|
||||||
|
});
|
||||||
|
|
||||||
|
it("keeps at least two inputs when disconnecting", () => {
|
||||||
|
const nodeType = createNodeType();
|
||||||
|
const node = createNode(["lora_stack1", "lora_stack2"]);
|
||||||
|
node.inputs[0].link = 11;
|
||||||
|
node.inputs[1].link = null; // slot 2 was just disconnected
|
||||||
|
node.onConnectionsChange = nodeType.prototype.onConnectionsChange;
|
||||||
|
|
||||||
|
node.onConnectionsChange(1, 1, false, makeLinkInfo());
|
||||||
|
|
||||||
|
expect(node.inputs.map((input) => input.name)).toEqual([
|
||||||
|
"lora_stack1",
|
||||||
|
"lora_stack2",
|
||||||
|
]);
|
||||||
|
expect(node.removeInput).not.toHaveBeenCalled();
|
||||||
|
});
|
||||||
|
|
||||||
|
it("does nothing while the graph is being configured", () => {
|
||||||
|
appMock.configuringGraph = true;
|
||||||
|
const nodeType = createNodeType();
|
||||||
|
const node = createNode(["lora_stack1", "lora_stack2"]);
|
||||||
|
node.onConnectionsChange = nodeType.prototype.onConnectionsChange;
|
||||||
|
|
||||||
|
node.onConnectionsChange(1, 1, true, makeLinkInfo());
|
||||||
|
|
||||||
|
expect(node.inputs.map((input) => input.name)).toEqual([
|
||||||
|
"lora_stack1",
|
||||||
|
"lora_stack2",
|
||||||
|
]);
|
||||||
|
expect(node.addInput).not.toHaveBeenCalled();
|
||||||
|
});
|
||||||
|
|
||||||
|
it("leaves legacy lora_stack_a/b inputs untouched", () => {
|
||||||
|
const nodeType = createNodeType();
|
||||||
|
const node = createNode(["lora_stack_a", "lora_stack_b"]);
|
||||||
|
node.onConnectionsChange = nodeType.prototype.onConnectionsChange;
|
||||||
|
|
||||||
|
node.onConnectionsChange(1, 0, true, makeLinkInfo());
|
||||||
|
|
||||||
|
expect(node.inputs.map((input) => input.name)).toEqual([
|
||||||
|
"lora_stack_a",
|
||||||
|
"lora_stack_b",
|
||||||
|
]);
|
||||||
|
expect(node.addInput).not.toHaveBeenCalled();
|
||||||
|
});
|
||||||
|
|
||||||
|
it("ensures two numbered inputs exist on creation", () => {
|
||||||
|
const node = createNode([]);
|
||||||
|
extension.nodeCreated(node, {});
|
||||||
|
|
||||||
|
expect(node.inputs.map((input) => input.name)).toEqual([
|
||||||
|
"lora_stack1",
|
||||||
|
"lora_stack2",
|
||||||
|
]);
|
||||||
|
});
|
||||||
|
|
||||||
|
it("does not add numbered inputs to legacy workflows", () => {
|
||||||
|
const node = createNode(["lora_stack_a", "lora_stack_b"]);
|
||||||
|
extension.nodeCreated(node, {});
|
||||||
|
|
||||||
|
expect(node.inputs.map((input) => input.name)).toEqual([
|
||||||
|
"lora_stack_a",
|
||||||
|
"lora_stack_b",
|
||||||
|
]);
|
||||||
|
});
|
||||||
|
});
|
||||||
@@ -0,0 +1,133 @@
|
|||||||
|
import { describe, it, beforeEach, expect, vi } from 'vitest';
|
||||||
|
|
||||||
|
const showToastMock = vi.fn();
|
||||||
|
const translateMock = vi.fn((key, params, fallback) => (typeof fallback === 'string' ? fallback : key));
|
||||||
|
const getNSFWLevelNameMock = vi.fn((level) => {
|
||||||
|
if (level >= 16) return 'XXX';
|
||||||
|
if (level >= 8) return 'X';
|
||||||
|
if (level >= 4) return 'R';
|
||||||
|
if (level >= 2) return 'PG13';
|
||||||
|
if (level >= 1) return 'PG';
|
||||||
|
return 'Unknown';
|
||||||
|
});
|
||||||
|
|
||||||
|
const loadingManagerStub = {
|
||||||
|
showSimpleLoading: vi.fn(),
|
||||||
|
showCancelButton: vi.fn(),
|
||||||
|
hide: vi.fn(),
|
||||||
|
};
|
||||||
|
|
||||||
|
const stateStub = {
|
||||||
|
currentPageType: 'recipes',
|
||||||
|
bulkMode: false,
|
||||||
|
selectedModels: new Set(),
|
||||||
|
loadingManager: loadingManagerStub,
|
||||||
|
virtualScroller: { updateSingleItem: vi.fn() },
|
||||||
|
global: { settings: {} },
|
||||||
|
};
|
||||||
|
|
||||||
|
const saveModelMetadataMock = vi.fn();
|
||||||
|
const getModelApiClientMock = vi.fn(() => ({ saveModelMetadata: saveModelMetadataMock }));
|
||||||
|
const updateRecipeMetadataMock = vi.fn(() => Promise.resolve({ success: true }));
|
||||||
|
|
||||||
|
vi.mock('../../../static/js/state/index.js', () => ({
|
||||||
|
state: stateStub,
|
||||||
|
getCurrentPageState: vi.fn(),
|
||||||
|
}));
|
||||||
|
|
||||||
|
vi.mock('../../../static/js/utils/uiHelpers.js', () => ({
|
||||||
|
showToast: showToastMock,
|
||||||
|
copyToClipboard: vi.fn(),
|
||||||
|
sendLoraToWorkflow: vi.fn(),
|
||||||
|
sendEmbeddingToWorkflow: vi.fn(),
|
||||||
|
buildLoraSyntax: vi.fn(),
|
||||||
|
getNSFWLevelName: getNSFWLevelNameMock,
|
||||||
|
}));
|
||||||
|
|
||||||
|
vi.mock('../../../static/js/api/modelApiFactory.js', () => ({
|
||||||
|
getModelApiClient: getModelApiClientMock,
|
||||||
|
resetAndReload: vi.fn(),
|
||||||
|
}));
|
||||||
|
|
||||||
|
vi.mock('../../../static/js/api/recipeApi.js', () => ({
|
||||||
|
RecipeSidebarApiClient: class {},
|
||||||
|
updateRecipeMetadata: updateRecipeMetadataMock,
|
||||||
|
extractRecipeId: vi.fn(),
|
||||||
|
}));
|
||||||
|
|
||||||
|
vi.mock('../../../static/js/api/apiConfig.js', () => ({
|
||||||
|
MODEL_TYPES: { LORA: 'loras', CHECKPOINT: 'checkpoints', EMBEDDING: 'embeddings' },
|
||||||
|
MODEL_CONFIG: {},
|
||||||
|
}));
|
||||||
|
|
||||||
|
vi.mock('../../../static/js/managers/ModalManager.js', () => ({
|
||||||
|
modalManager: { showModal: vi.fn(), closeModal: vi.fn() },
|
||||||
|
}));
|
||||||
|
|
||||||
|
vi.mock('../../../static/js/components/shared/ModelCard.js', () => ({
|
||||||
|
updateCardsForBulkMode: vi.fn(),
|
||||||
|
}));
|
||||||
|
|
||||||
|
vi.mock('../../../static/js/utils/i18nHelpers.js', () => ({
|
||||||
|
translate: translateMock,
|
||||||
|
}));
|
||||||
|
|
||||||
|
vi.mock('../../../static/js/utils/priorityTagHelpers.js', () => ({
|
||||||
|
getPriorityTagSuggestions: vi.fn(),
|
||||||
|
}));
|
||||||
|
|
||||||
|
vi.mock('../../../static/js/components/shared/NsfwLevelSelector.js', () => ({
|
||||||
|
getNsfwLevelSelector: vi.fn(),
|
||||||
|
}));
|
||||||
|
|
||||||
|
describe('BulkManager bulk content rating', () => {
|
||||||
|
beforeEach(() => {
|
||||||
|
vi.clearAllMocks();
|
||||||
|
stateStub.currentPageType = 'recipes';
|
||||||
|
stateStub.bulkMode = false;
|
||||||
|
stateStub.selectedModels.clear();
|
||||||
|
saveModelMetadataMock.mockResolvedValue(undefined);
|
||||||
|
updateRecipeMetadataMock.mockResolvedValue({ success: true });
|
||||||
|
});
|
||||||
|
|
||||||
|
async function createBulkManager() {
|
||||||
|
const { BulkManager } = await import('../../../static/js/managers/BulkManager.js');
|
||||||
|
return new BulkManager();
|
||||||
|
}
|
||||||
|
|
||||||
|
it('exposes the content rating action on the recipes page action config', async () => {
|
||||||
|
const bulk = await createBulkManager();
|
||||||
|
expect(bulk.actionConfig.recipes.setContentRating).toBe(true);
|
||||||
|
});
|
||||||
|
|
||||||
|
it('persists the rating through the recipe API when on the recipes page', async () => {
|
||||||
|
const bulk = await createBulkManager();
|
||||||
|
stateStub.currentPageType = 'recipes';
|
||||||
|
stateStub.selectedModels.add('/recipes/test.webp');
|
||||||
|
|
||||||
|
const ok = await bulk.setBulkContentRating(4, ['/recipes/test.webp']);
|
||||||
|
|
||||||
|
expect(ok).toBe(true);
|
||||||
|
expect(updateRecipeMetadataMock).toHaveBeenCalledWith('/recipes/test.webp', { preview_nsfw_level: 4 });
|
||||||
|
expect(updateRecipeMetadataMock).toHaveBeenCalledTimes(1);
|
||||||
|
expect(saveModelMetadataMock).not.toHaveBeenCalled();
|
||||||
|
expect(showToastMock).toHaveBeenCalledWith(
|
||||||
|
'toast.models.bulkContentRatingSet',
|
||||||
|
{ count: 1, level: 'R' },
|
||||||
|
'success'
|
||||||
|
);
|
||||||
|
});
|
||||||
|
|
||||||
|
it('persists the rating through the model API on model pages', async () => {
|
||||||
|
const bulk = await createBulkManager();
|
||||||
|
stateStub.currentPageType = 'loras';
|
||||||
|
stateStub.selectedModels.add('/models/test.safetensors');
|
||||||
|
|
||||||
|
const ok = await bulk.setBulkContentRating(8, ['/models/test.safetensors']);
|
||||||
|
|
||||||
|
expect(ok).toBe(true);
|
||||||
|
expect(saveModelMetadataMock).toHaveBeenCalledWith('/models/test.safetensors', { preview_nsfw_level: 8 });
|
||||||
|
expect(saveModelMetadataMock).toHaveBeenCalledTimes(1);
|
||||||
|
expect(updateRecipeMetadataMock).not.toHaveBeenCalled();
|
||||||
|
});
|
||||||
|
});
|
||||||
@@ -0,0 +1,434 @@
|
|||||||
|
import { afterEach, beforeEach, describe, expect, it, vi } from 'vitest';
|
||||||
|
|
||||||
|
const {
|
||||||
|
DOWNLOAD_MANAGER_MODULE,
|
||||||
|
MODAL_MANAGER_MODULE,
|
||||||
|
UI_HELPERS_MODULE,
|
||||||
|
STATE_MODULE,
|
||||||
|
LOADING_MANAGER_MODULE,
|
||||||
|
API_FACTORY_MODULE,
|
||||||
|
STORAGE_HELPERS_MODULE,
|
||||||
|
FOLDER_TREE_MANAGER_MODULE,
|
||||||
|
I18N_HELPERS_MODULE,
|
||||||
|
SUMMARY_MODULE,
|
||||||
|
mockApiClient,
|
||||||
|
mockLoadingManager,
|
||||||
|
showToastMock,
|
||||||
|
showDownloadBatchSummaryMock,
|
||||||
|
resetAndReloadMock,
|
||||||
|
} = vi.hoisted(() => {
|
||||||
|
// Shared API client returned by the mocked getModelApiClient factory.
|
||||||
|
const mockApiClient = {
|
||||||
|
apiConfig: {
|
||||||
|
config: {
|
||||||
|
displayName: 'LoRA',
|
||||||
|
singularName: 'lora',
|
||||||
|
},
|
||||||
|
},
|
||||||
|
downloadModel: vi.fn(),
|
||||||
|
downloadHfModel: vi.fn(),
|
||||||
|
cancelDownload: vi.fn(),
|
||||||
|
getPageState: vi.fn(() => ({})),
|
||||||
|
};
|
||||||
|
|
||||||
|
// Shared loading manager served both via state.loadingManager and the
|
||||||
|
// LoadingManager constructor mock.
|
||||||
|
const mockLoadingManager = {
|
||||||
|
showSimpleLoading: vi.fn(),
|
||||||
|
hide: vi.fn(),
|
||||||
|
restoreProgressBar: vi.fn(),
|
||||||
|
showDownloadProgress: vi.fn(() => vi.fn()),
|
||||||
|
setStatus: vi.fn(),
|
||||||
|
showCancelButton: vi.fn(),
|
||||||
|
};
|
||||||
|
|
||||||
|
return {
|
||||||
|
DOWNLOAD_MANAGER_MODULE: new URL('../../../static/js/managers/DownloadManager.js', import.meta.url).pathname,
|
||||||
|
MODAL_MANAGER_MODULE: new URL('../../../static/js/managers/ModalManager.js', import.meta.url).pathname,
|
||||||
|
UI_HELPERS_MODULE: new URL('../../../static/js/utils/uiHelpers.js', import.meta.url).pathname,
|
||||||
|
STATE_MODULE: new URL('../../../static/js/state/index.js', import.meta.url).pathname,
|
||||||
|
LOADING_MANAGER_MODULE: new URL('../../../static/js/managers/LoadingManager.js', import.meta.url).pathname,
|
||||||
|
API_FACTORY_MODULE: new URL('../../../static/js/api/modelApiFactory.js', import.meta.url).pathname,
|
||||||
|
STORAGE_HELPERS_MODULE: new URL('../../../static/js/utils/storageHelpers.js', import.meta.url).pathname,
|
||||||
|
FOLDER_TREE_MANAGER_MODULE: new URL('../../../static/js/components/FolderTreeManager.js', import.meta.url).pathname,
|
||||||
|
I18N_HELPERS_MODULE: new URL('../../../static/js/utils/i18nHelpers.js', import.meta.url).pathname,
|
||||||
|
SUMMARY_MODULE: new URL('../../../static/js/components/DownloadBatchSummaryModal.js', import.meta.url).pathname,
|
||||||
|
mockApiClient,
|
||||||
|
mockLoadingManager,
|
||||||
|
showToastMock: vi.fn(),
|
||||||
|
showDownloadBatchSummaryMock: vi.fn(),
|
||||||
|
resetAndReloadMock: vi.fn(),
|
||||||
|
};
|
||||||
|
});
|
||||||
|
|
||||||
|
vi.mock(MODAL_MANAGER_MODULE, () => ({
|
||||||
|
modalManager: {
|
||||||
|
showModal: vi.fn(),
|
||||||
|
closeModal: vi.fn(),
|
||||||
|
},
|
||||||
|
}));
|
||||||
|
|
||||||
|
vi.mock(UI_HELPERS_MODULE, () => ({
|
||||||
|
showToast: showToastMock,
|
||||||
|
}));
|
||||||
|
|
||||||
|
vi.mock(STATE_MODULE, () => ({
|
||||||
|
state: {
|
||||||
|
global: {
|
||||||
|
settings: {},
|
||||||
|
},
|
||||||
|
loadingManager: mockLoadingManager,
|
||||||
|
},
|
||||||
|
}));
|
||||||
|
|
||||||
|
vi.mock(LOADING_MANAGER_MODULE, () => ({
|
||||||
|
LoadingManager: vi.fn(() => mockLoadingManager),
|
||||||
|
}));
|
||||||
|
|
||||||
|
vi.mock(API_FACTORY_MODULE, () => ({
|
||||||
|
getModelApiClient: vi.fn(() => mockApiClient),
|
||||||
|
resetAndReload: resetAndReloadMock,
|
||||||
|
}));
|
||||||
|
|
||||||
|
vi.mock(STORAGE_HELPERS_MODULE, () => ({
|
||||||
|
getStorageItem: vi.fn((_key, defaultValue) => defaultValue),
|
||||||
|
setStorageItem: vi.fn(),
|
||||||
|
}));
|
||||||
|
|
||||||
|
vi.mock(FOLDER_TREE_MANAGER_MODULE, () => ({
|
||||||
|
FolderTreeManager: vi.fn(() => ({
|
||||||
|
clearSelection: vi.fn(),
|
||||||
|
init: vi.fn(),
|
||||||
|
})),
|
||||||
|
}));
|
||||||
|
|
||||||
|
vi.mock(I18N_HELPERS_MODULE, () => ({
|
||||||
|
translate: vi.fn((_, __, fallback) => fallback ?? ''),
|
||||||
|
}));
|
||||||
|
|
||||||
|
vi.mock(SUMMARY_MODULE, () => ({
|
||||||
|
showDownloadBatchSummary: showDownloadBatchSummaryMock,
|
||||||
|
}));
|
||||||
|
|
||||||
|
/**
|
||||||
|
* Fake WebSocket used by executeBatchDownload. Resolves `onopen` on the
|
||||||
|
* microtask queue right after construction (which happens after the real
|
||||||
|
* code has assigned `onopen`), so the open promise resolves deterministically
|
||||||
|
* without real timers.
|
||||||
|
*/
|
||||||
|
class FakeWebSocket {
|
||||||
|
static instances = [];
|
||||||
|
|
||||||
|
constructor(url) {
|
||||||
|
this.url = url;
|
||||||
|
this.onopen = null;
|
||||||
|
this.onmessage = null;
|
||||||
|
this.onerror = null;
|
||||||
|
this.close = vi.fn();
|
||||||
|
FakeWebSocket.instances.push(this);
|
||||||
|
queueMicrotask(() => {
|
||||||
|
if (this.onopen) this.onopen();
|
||||||
|
});
|
||||||
|
}
|
||||||
|
|
||||||
|
static get lastInstance() {
|
||||||
|
return FakeWebSocket.instances[FakeWebSocket.instances.length - 1];
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
describe('DownloadManager batch download summary flow', () => {
|
||||||
|
let DownloadManager;
|
||||||
|
let manager;
|
||||||
|
|
||||||
|
const options = { modelRoot: '/models/loras', targetFolder: '', useDefaultPaths: true };
|
||||||
|
|
||||||
|
const makeItem = (modelId, versionId, name) => ({
|
||||||
|
modelId,
|
||||||
|
displayName: name,
|
||||||
|
selectedVersion: { id: versionId, name, existsLocally: false },
|
||||||
|
});
|
||||||
|
|
||||||
|
const item0 = makeItem('111', 'v1', 'Model A');
|
||||||
|
const item1 = makeItem('222', 'v2', 'Model B');
|
||||||
|
|
||||||
|
beforeEach(async () => {
|
||||||
|
document.body.innerHTML = '';
|
||||||
|
FakeWebSocket.instances = [];
|
||||||
|
|
||||||
|
// Reset the shared mocks so mockResolvedValueOnce queues and call
|
||||||
|
// history never leak between tests.
|
||||||
|
mockApiClient.downloadModel.mockReset();
|
||||||
|
mockApiClient.downloadHfModel.mockReset();
|
||||||
|
mockApiClient.cancelDownload.mockReset();
|
||||||
|
showToastMock.mockClear();
|
||||||
|
showDownloadBatchSummaryMock.mockClear();
|
||||||
|
resetAndReloadMock.mockClear();
|
||||||
|
mockLoadingManager.hide.mockClear();
|
||||||
|
mockLoadingManager.setStatus.mockClear();
|
||||||
|
mockLoadingManager.showCancelButton.mockClear();
|
||||||
|
mockLoadingManager.showDownloadProgress.mockClear();
|
||||||
|
|
||||||
|
vi.stubGlobal('WebSocket', FakeWebSocket);
|
||||||
|
vi.resetModules();
|
||||||
|
({ DownloadManager } = await import(DOWNLOAD_MANAGER_MODULE));
|
||||||
|
manager = new DownloadManager();
|
||||||
|
// The constructor leaves apiClient null; executeBatchDownload reads it
|
||||||
|
// directly, so point it at the shared mocked client.
|
||||||
|
manager.apiClient = mockApiClient;
|
||||||
|
});
|
||||||
|
|
||||||
|
afterEach(() => {
|
||||||
|
document.body.innerHTML = '';
|
||||||
|
vi.unstubAllGlobals();
|
||||||
|
});
|
||||||
|
|
||||||
|
it('shows the success toast when every item downloads successfully', async () => {
|
||||||
|
mockApiClient.downloadModel.mockResolvedValue({ success: true });
|
||||||
|
|
||||||
|
await manager.executeBatchDownload([item0, item1], options);
|
||||||
|
|
||||||
|
expect(mockApiClient.downloadModel).toHaveBeenCalledTimes(2);
|
||||||
|
// Each item is downloaded with its own modelId + versionId.
|
||||||
|
expect(mockApiClient.downloadModel.mock.calls[0][0]).toBe('111');
|
||||||
|
expect(mockApiClient.downloadModel.mock.calls[0][1]).toBe('v1');
|
||||||
|
expect(mockApiClient.downloadModel.mock.calls[1][0]).toBe('222');
|
||||||
|
expect(mockApiClient.downloadModel.mock.calls[1][1]).toBe('v2');
|
||||||
|
|
||||||
|
expect(showDownloadBatchSummaryMock).not.toHaveBeenCalled();
|
||||||
|
expect(showToastMock).toHaveBeenCalledTimes(1);
|
||||||
|
expect(showToastMock).toHaveBeenCalledWith('toast.loras.allDownloadSuccessful', { count: 2 }, 'success');
|
||||||
|
expect(resetAndReloadMock).toHaveBeenCalledWith(true);
|
||||||
|
});
|
||||||
|
|
||||||
|
it('shows a partial-failure summary when some items fail', async () => {
|
||||||
|
// The failing item has no displayName/filename, so the resolved entry
|
||||||
|
// name falls back to the selected version name.
|
||||||
|
const unnamedItem = { modelId: '333', selectedVersion: { id: 'v3', name: 'V3', existsLocally: false } };
|
||||||
|
mockApiClient.downloadModel
|
||||||
|
.mockResolvedValueOnce({ success: false, error: 'rate limited' })
|
||||||
|
.mockResolvedValueOnce({ success: true });
|
||||||
|
|
||||||
|
await manager.executeBatchDownload([unnamedItem, item1], options);
|
||||||
|
|
||||||
|
expect(showDownloadBatchSummaryMock).toHaveBeenCalledTimes(1);
|
||||||
|
const summary = showDownloadBatchSummaryMock.mock.calls[0][0];
|
||||||
|
expect(summary.total).toBe(2);
|
||||||
|
expect(summary.completed).toBe(1);
|
||||||
|
expect(summary.failedItems).toHaveLength(1);
|
||||||
|
expect(summary.failedItems[0].item).toBe(unnamedItem);
|
||||||
|
expect(summary.failedItems[0].error).toBe('rate limited');
|
||||||
|
// The resolved display name is carried on the failed entry.
|
||||||
|
expect(summary.failedItems[0].name).toBe('V3');
|
||||||
|
expect(summary.onRetry).toEqual(expect.any(Function));
|
||||||
|
|
||||||
|
// No success toast and no downloadPartialSuccess toast for this path.
|
||||||
|
expect(showToastMock).not.toHaveBeenCalledWith('toast.loras.allDownloadSuccessful', expect.anything(), 'success');
|
||||||
|
expect(showToastMock).not.toHaveBeenCalledWith('toast.loras.downloadPartialSuccess', expect.anything(), expect.anything());
|
||||||
|
});
|
||||||
|
|
||||||
|
it('shows an all-failed summary when every item fails', async () => {
|
||||||
|
mockApiClient.downloadModel.mockResolvedValue({ success: false, error: 'x' });
|
||||||
|
|
||||||
|
await manager.executeBatchDownload([item0, item1], options);
|
||||||
|
|
||||||
|
expect(showDownloadBatchSummaryMock).toHaveBeenCalledTimes(1);
|
||||||
|
const summary = showDownloadBatchSummaryMock.mock.calls[0][0];
|
||||||
|
expect(summary.total).toBe(2);
|
||||||
|
expect(summary.completed).toBe(0);
|
||||||
|
expect(summary.failedItems).toHaveLength(2);
|
||||||
|
expect(summary.failedItems[0].item).toBe(item0);
|
||||||
|
expect(summary.failedItems[1].item).toBe(item1);
|
||||||
|
expect(showToastMock).not.toHaveBeenCalledWith('toast.loras.allDownloadSuccessful', expect.anything(), expect.anything());
|
||||||
|
});
|
||||||
|
|
||||||
|
it('records the error message when downloadModel rejects', async () => {
|
||||||
|
// The item carries a filename but no displayName, so the resolved entry
|
||||||
|
// name comes from the filename.
|
||||||
|
const filenameItem = { modelId: '444', filename: 'model.safetensors', selectedVersion: { id: 'v4' } };
|
||||||
|
mockApiClient.downloadModel.mockRejectedValue(new Error('network down'));
|
||||||
|
|
||||||
|
await manager.executeBatchDownload([filenameItem], options);
|
||||||
|
|
||||||
|
expect(showDownloadBatchSummaryMock).toHaveBeenCalledTimes(1);
|
||||||
|
const summary = showDownloadBatchSummaryMock.mock.calls[0][0];
|
||||||
|
expect(summary.total).toBe(1);
|
||||||
|
expect(summary.completed).toBe(0);
|
||||||
|
expect(summary.failedItems).toHaveLength(1);
|
||||||
|
expect(summary.failedItems[0].item).toBe(filenameItem);
|
||||||
|
expect(summary.failedItems[0].error).toBe('network down');
|
||||||
|
expect(summary.failedItems[0].name).toBe('model.safetensors');
|
||||||
|
});
|
||||||
|
|
||||||
|
it('retries the failed subset through onRetry with unwrapped items', async () => {
|
||||||
|
mockApiClient.downloadModel
|
||||||
|
.mockResolvedValueOnce({ success: false, error: 'rate limited' })
|
||||||
|
.mockResolvedValueOnce({ success: true });
|
||||||
|
|
||||||
|
await manager.executeBatchDownload([item0, item1], options);
|
||||||
|
|
||||||
|
expect(showDownloadBatchSummaryMock).toHaveBeenCalledTimes(1);
|
||||||
|
const summary = showDownloadBatchSummaryMock.mock.calls[0][0];
|
||||||
|
expect(summary.failedItems).toHaveLength(1);
|
||||||
|
|
||||||
|
// Retry the exact failed subset returned by the summary. The onRetry
|
||||||
|
// callback unwraps the { item, error } entries back into raw model items
|
||||||
|
// before re-running executeBatchDownload. Make the retried item fail
|
||||||
|
// again so a second summary is produced.
|
||||||
|
mockApiClient.downloadModel.mockResolvedValueOnce({ success: false, error: 'still rate limited' });
|
||||||
|
await summary.onRetry(summary.failedItems);
|
||||||
|
|
||||||
|
// downloadModel is called a third time — only for the failed item (item0),
|
||||||
|
// NOT for the item that already succeeded (item1).
|
||||||
|
expect(mockApiClient.downloadModel).toHaveBeenCalledTimes(3);
|
||||||
|
const retryCall = mockApiClient.downloadModel.mock.calls[2];
|
||||||
|
expect(retryCall[0]).toBe(item0.modelId);
|
||||||
|
expect(retryCall[1]).toBe(item0.selectedVersion.id);
|
||||||
|
|
||||||
|
// A fresh summary is produced for the retry run (call count 1 -> 2).
|
||||||
|
expect(showDownloadBatchSummaryMock).toHaveBeenCalledTimes(2);
|
||||||
|
const retrySummary = showDownloadBatchSummaryMock.mock.calls[1][0];
|
||||||
|
expect(retrySummary.total).toBe(1);
|
||||||
|
expect(retrySummary.completed).toBe(0);
|
||||||
|
expect(retrySummary.failedItems).toHaveLength(1);
|
||||||
|
expect(retrySummary.failedItems[0].item).toBe(item0);
|
||||||
|
expect(retrySummary.failedItems[0].error).toBe('still rate limited');
|
||||||
|
});
|
||||||
|
|
||||||
|
it('stops the batch without showing a summary when cancelled before downloads start', async () => {
|
||||||
|
const downloadPromise = manager.executeBatchDownload([item0, item1], options);
|
||||||
|
|
||||||
|
// showCancelButton captured the cancel callback synchronously. Invoking it
|
||||||
|
// sets `cancelled = true` before the download loop runs (the loop only
|
||||||
|
// starts after the WebSocket open promise resolves on the microtask queue).
|
||||||
|
const cancelCallback = mockLoadingManager.showCancelButton.mock.calls[0][0];
|
||||||
|
const cancelPromise = cancelCallback();
|
||||||
|
|
||||||
|
await Promise.all([downloadPromise, cancelPromise]);
|
||||||
|
|
||||||
|
expect(mockApiClient.downloadModel).not.toHaveBeenCalled();
|
||||||
|
expect(showDownloadBatchSummaryMock).not.toHaveBeenCalled();
|
||||||
|
expect(showToastMock).toHaveBeenCalledWith(
|
||||||
|
'toast.downloads.downloadStopped',
|
||||||
|
expect.anything(),
|
||||||
|
'info',
|
||||||
|
expect.stringContaining('Download cancelled')
|
||||||
|
);
|
||||||
|
expect(resetAndReloadMock).toHaveBeenCalledWith(true);
|
||||||
|
});
|
||||||
|
|
||||||
|
it('shows the batch summary for a single CivitAI download resolved as a failure', async () => {
|
||||||
|
mockApiClient.downloadModel.mockResolvedValue({ success: false, error: 'rate limited' });
|
||||||
|
|
||||||
|
await manager.executeDownloadWithProgress({
|
||||||
|
modelId: '111',
|
||||||
|
versionId: 'v1',
|
||||||
|
versionName: 'V1',
|
||||||
|
modelRoot: '/m',
|
||||||
|
useDefaultPaths: true,
|
||||||
|
});
|
||||||
|
|
||||||
|
expect(showDownloadBatchSummaryMock).toHaveBeenCalledTimes(1);
|
||||||
|
const summary = showDownloadBatchSummaryMock.mock.calls[0][0];
|
||||||
|
expect(summary.total).toBe(1);
|
||||||
|
expect(summary.completed).toBe(0);
|
||||||
|
expect(summary.failedItems).toHaveLength(1);
|
||||||
|
expect(summary.failedItems[0].item.modelId).toBe('111');
|
||||||
|
expect(summary.failedItems[0].item.versionId).toBe('v1');
|
||||||
|
expect(summary.failedItems[0].item.url).toEqual(expect.stringContaining('civitai.com/models/111'));
|
||||||
|
expect(summary.failedItems[0].error).toBe('rate limited');
|
||||||
|
expect(summary.failedItems[0].name).toBe('V1');
|
||||||
|
expect(summary.onRetry).toEqual(expect.any(Function));
|
||||||
|
expect(showToastMock).not.toHaveBeenCalledWith('toast.loras.downloadCompleted', expect.anything(), 'success');
|
||||||
|
});
|
||||||
|
|
||||||
|
it('shows the batch summary when a single download throws', async () => {
|
||||||
|
mockApiClient.downloadModel.mockRejectedValue(new Error('network down'));
|
||||||
|
|
||||||
|
await manager.executeDownloadWithProgress({
|
||||||
|
modelId: '111',
|
||||||
|
versionId: 'v1',
|
||||||
|
versionName: 'V1',
|
||||||
|
modelRoot: '/m',
|
||||||
|
useDefaultPaths: true,
|
||||||
|
});
|
||||||
|
|
||||||
|
expect(showDownloadBatchSummaryMock).toHaveBeenCalledTimes(1);
|
||||||
|
const summary = showDownloadBatchSummaryMock.mock.calls[0][0];
|
||||||
|
expect(summary.total).toBe(1);
|
||||||
|
expect(summary.completed).toBe(0);
|
||||||
|
expect(summary.failedItems).toHaveLength(1);
|
||||||
|
expect(summary.failedItems[0].error).toBe('network down');
|
||||||
|
expect(summary.failedItems[0].item.url).toEqual(expect.stringContaining('civitai.com/models/111'));
|
||||||
|
expect(showToastMock).not.toHaveBeenCalledWith('toast.loras.downloadCompleted', expect.anything(), 'success');
|
||||||
|
});
|
||||||
|
|
||||||
|
it('keeps the success toast and skips the summary for a successful single download', async () => {
|
||||||
|
mockApiClient.downloadModel.mockResolvedValue({ success: true });
|
||||||
|
|
||||||
|
const result = await manager.executeDownloadWithProgress({
|
||||||
|
modelId: '111',
|
||||||
|
versionId: 'v1',
|
||||||
|
versionName: 'V1',
|
||||||
|
modelRoot: '/m',
|
||||||
|
useDefaultPaths: true,
|
||||||
|
});
|
||||||
|
|
||||||
|
expect(result).toBe(true);
|
||||||
|
expect(showDownloadBatchSummaryMock).not.toHaveBeenCalled();
|
||||||
|
expect(showToastMock).toHaveBeenCalledWith('toast.loras.downloadCompleted', {}, 'success');
|
||||||
|
expect(resetAndReloadMock).toHaveBeenCalledWith(true);
|
||||||
|
});
|
||||||
|
|
||||||
|
it('retries a failed single download through onRetry with the same params', async () => {
|
||||||
|
mockApiClient.downloadModel
|
||||||
|
.mockResolvedValueOnce({ success: false, error: 'rate limited' })
|
||||||
|
.mockResolvedValueOnce({ success: true });
|
||||||
|
|
||||||
|
await manager.executeDownloadWithProgress({
|
||||||
|
modelId: '111',
|
||||||
|
versionId: 'v1',
|
||||||
|
versionName: 'V1',
|
||||||
|
modelRoot: '/m',
|
||||||
|
useDefaultPaths: true,
|
||||||
|
});
|
||||||
|
|
||||||
|
expect(showDownloadBatchSummaryMock).toHaveBeenCalledTimes(1);
|
||||||
|
const summary = showDownloadBatchSummaryMock.mock.calls[0][0];
|
||||||
|
expect(summary.failedItems).toHaveLength(1);
|
||||||
|
|
||||||
|
await summary.onRetry();
|
||||||
|
|
||||||
|
expect(mockApiClient.downloadModel).toHaveBeenCalledTimes(2);
|
||||||
|
const retryCall = mockApiClient.downloadModel.mock.calls[1];
|
||||||
|
expect(retryCall[0]).toBe('111');
|
||||||
|
expect(retryCall[1]).toBe('v1');
|
||||||
|
expect(showDownloadBatchSummaryMock).toHaveBeenCalledTimes(1);
|
||||||
|
expect(showToastMock).toHaveBeenCalledWith('toast.loras.downloadCompleted', {}, 'success');
|
||||||
|
});
|
||||||
|
|
||||||
|
it('shows a summary for HF partial failure and retries only the failed files', async () => {
|
||||||
|
manager.hfRepoId = 'user/repo';
|
||||||
|
manager.hfSelectedFiles = ['a.safetensors', 'b.safetensors'];
|
||||||
|
mockApiClient.downloadHfModel
|
||||||
|
.mockResolvedValueOnce({ success: true })
|
||||||
|
.mockResolvedValueOnce({ success: false, error: 'denied' });
|
||||||
|
|
||||||
|
const result = await manager._downloadHfSingle({ modelRoot: '/m', useDefaultPaths: true });
|
||||||
|
|
||||||
|
expect(result).toBe(false);
|
||||||
|
expect(showDownloadBatchSummaryMock).toHaveBeenCalledTimes(1);
|
||||||
|
const summary = showDownloadBatchSummaryMock.mock.calls[0][0];
|
||||||
|
expect(summary.total).toBe(2);
|
||||||
|
expect(summary.completed).toBe(1);
|
||||||
|
expect(summary.failedItems).toHaveLength(1);
|
||||||
|
expect(summary.failedItems[0].name).toBe('b.safetensors');
|
||||||
|
expect(summary.failedItems[0].item.url).toEqual(
|
||||||
|
expect.stringContaining('huggingface.co/user/repo/blob/main/b.safetensors')
|
||||||
|
);
|
||||||
|
|
||||||
|
await summary.onRetry();
|
||||||
|
|
||||||
|
expect(mockApiClient.downloadHfModel).toHaveBeenCalledTimes(3);
|
||||||
|
expect(mockApiClient.downloadHfModel.mock.calls[2][0].filename).toBe('b.safetensors');
|
||||||
|
});
|
||||||
|
});
|
||||||
@@ -0,0 +1,264 @@
|
|||||||
|
import { afterEach, beforeEach, describe, expect, it, vi } from 'vitest';
|
||||||
|
|
||||||
|
const {
|
||||||
|
DOWNLOAD_MANAGER_MODULE,
|
||||||
|
MODAL_MANAGER_MODULE,
|
||||||
|
UI_HELPERS_MODULE,
|
||||||
|
STATE_MODULE,
|
||||||
|
LOADING_MANAGER_MODULE,
|
||||||
|
API_FACTORY_MODULE,
|
||||||
|
STORAGE_HELPERS_MODULE,
|
||||||
|
FOLDER_TREE_MANAGER_MODULE,
|
||||||
|
I18N_HELPERS_MODULE,
|
||||||
|
SUMMARY_MODULE,
|
||||||
|
mockApiClient,
|
||||||
|
mockLoadingManager,
|
||||||
|
showToastMock,
|
||||||
|
showDownloadBatchSummaryMock,
|
||||||
|
resetAndReloadMock,
|
||||||
|
} = vi.hoisted(() => {
|
||||||
|
// Shared API client returned by the mocked getModelApiClient factory.
|
||||||
|
const mockApiClient = {
|
||||||
|
apiConfig: {
|
||||||
|
config: {
|
||||||
|
displayName: 'LoRA',
|
||||||
|
singularName: 'lora',
|
||||||
|
},
|
||||||
|
},
|
||||||
|
fetchCivitaiVersions: vi.fn(),
|
||||||
|
downloadModel: vi.fn(),
|
||||||
|
downloadHfModel: vi.fn(),
|
||||||
|
cancelDownload: vi.fn(),
|
||||||
|
getPageState: vi.fn(() => ({})),
|
||||||
|
};
|
||||||
|
|
||||||
|
// Shared loading manager served both via state.loadingManager and the
|
||||||
|
// LoadingManager constructor mock.
|
||||||
|
const mockLoadingManager = {
|
||||||
|
showSimpleLoading: vi.fn(),
|
||||||
|
setStatus: vi.fn(),
|
||||||
|
hide: vi.fn(),
|
||||||
|
restoreProgressBar: vi.fn(),
|
||||||
|
showDownloadProgress: vi.fn(() => vi.fn()),
|
||||||
|
showCancelButton: vi.fn(),
|
||||||
|
};
|
||||||
|
|
||||||
|
return {
|
||||||
|
DOWNLOAD_MANAGER_MODULE: new URL('../../../static/js/managers/DownloadManager.js', import.meta.url).pathname,
|
||||||
|
MODAL_MANAGER_MODULE: new URL('../../../static/js/managers/ModalManager.js', import.meta.url).pathname,
|
||||||
|
UI_HELPERS_MODULE: new URL('../../../static/js/utils/uiHelpers.js', import.meta.url).pathname,
|
||||||
|
STATE_MODULE: new URL('../../../static/js/state/index.js', import.meta.url).pathname,
|
||||||
|
LOADING_MANAGER_MODULE: new URL('../../../static/js/managers/LoadingManager.js', import.meta.url).pathname,
|
||||||
|
API_FACTORY_MODULE: new URL('../../../static/js/api/modelApiFactory.js', import.meta.url).pathname,
|
||||||
|
STORAGE_HELPERS_MODULE: new URL('../../../static/js/utils/storageHelpers.js', import.meta.url).pathname,
|
||||||
|
FOLDER_TREE_MANAGER_MODULE: new URL('../../../static/js/components/FolderTreeManager.js', import.meta.url).pathname,
|
||||||
|
I18N_HELPERS_MODULE: new URL('../../../static/js/utils/i18nHelpers.js', import.meta.url).pathname,
|
||||||
|
SUMMARY_MODULE: new URL('../../../static/js/components/DownloadBatchSummaryModal.js', import.meta.url).pathname,
|
||||||
|
mockApiClient,
|
||||||
|
mockLoadingManager,
|
||||||
|
showToastMock: vi.fn(),
|
||||||
|
showDownloadBatchSummaryMock: vi.fn(),
|
||||||
|
resetAndReloadMock: vi.fn(),
|
||||||
|
};
|
||||||
|
});
|
||||||
|
|
||||||
|
vi.mock(MODAL_MANAGER_MODULE, () => ({
|
||||||
|
modalManager: {
|
||||||
|
showModal: vi.fn(),
|
||||||
|
closeModal: vi.fn(),
|
||||||
|
},
|
||||||
|
}));
|
||||||
|
|
||||||
|
vi.mock(UI_HELPERS_MODULE, () => ({
|
||||||
|
showToast: showToastMock,
|
||||||
|
}));
|
||||||
|
|
||||||
|
vi.mock(STATE_MODULE, () => ({
|
||||||
|
state: {
|
||||||
|
global: {
|
||||||
|
settings: {},
|
||||||
|
},
|
||||||
|
loadingManager: mockLoadingManager,
|
||||||
|
},
|
||||||
|
}));
|
||||||
|
|
||||||
|
vi.mock(LOADING_MANAGER_MODULE, () => ({
|
||||||
|
LoadingManager: vi.fn(() => mockLoadingManager),
|
||||||
|
}));
|
||||||
|
|
||||||
|
vi.mock(API_FACTORY_MODULE, () => ({
|
||||||
|
getModelApiClient: vi.fn(() => mockApiClient),
|
||||||
|
resetAndReload: resetAndReloadMock,
|
||||||
|
}));
|
||||||
|
|
||||||
|
vi.mock(STORAGE_HELPERS_MODULE, () => ({
|
||||||
|
getStorageItem: vi.fn((_key, defaultValue) => defaultValue),
|
||||||
|
setStorageItem: vi.fn(),
|
||||||
|
}));
|
||||||
|
|
||||||
|
vi.mock(FOLDER_TREE_MANAGER_MODULE, () => ({
|
||||||
|
FolderTreeManager: vi.fn(() => ({
|
||||||
|
clearSelection: vi.fn(),
|
||||||
|
init: vi.fn(),
|
||||||
|
})),
|
||||||
|
}));
|
||||||
|
|
||||||
|
vi.mock(I18N_HELPERS_MODULE, () => ({
|
||||||
|
translate: vi.fn((_, __, fallback) => fallback ?? ''),
|
||||||
|
}));
|
||||||
|
|
||||||
|
vi.mock(SUMMARY_MODULE, () => ({
|
||||||
|
showDownloadBatchSummary: showDownloadBatchSummaryMock,
|
||||||
|
}));
|
||||||
|
|
||||||
|
/** Minimal DOM used by validateAndFetchVersions / showVersionStep / batch preview. */
|
||||||
|
function setupDownloadDom() {
|
||||||
|
document.body.innerHTML = `
|
||||||
|
<div id="downloadModal">
|
||||||
|
<div class="download-step" id="urlStep"></div>
|
||||||
|
<div class="download-step" id="versionStep"></div>
|
||||||
|
<div class="download-step" id="fileSelectionStep"></div>
|
||||||
|
<div class="download-step" id="downloadLocationStep"></div>
|
||||||
|
<div id="batchPreviewStep"></div>
|
||||||
|
<textarea id="modelUrl"></textarea>
|
||||||
|
<div id="urlError"></div>
|
||||||
|
<div id="versionList"></div>
|
||||||
|
<div id="fileSelectionList"></div>
|
||||||
|
<div id="fileSelectionVersionName"></div>
|
||||||
|
<button id="nextFromVersion"></button>
|
||||||
|
<button id="nextFromBatchBtn"></button>
|
||||||
|
<div id="downloadModalTitle"></div>
|
||||||
|
<div id="batchPreviewList"></div>
|
||||||
|
<div id="modelRoot"></div>
|
||||||
|
<div id="folderPath"></div>
|
||||||
|
<div id="targetPathDisplay"></div>
|
||||||
|
<input id="useDefaultPath" />
|
||||||
|
<div id="manualPathSelection"></div>
|
||||||
|
</div>
|
||||||
|
`;
|
||||||
|
}
|
||||||
|
|
||||||
|
function makeVersion(id, name) {
|
||||||
|
return {
|
||||||
|
id,
|
||||||
|
name,
|
||||||
|
baseModel: 'SDXL',
|
||||||
|
createdAt: '2025-01-01T00:00:00Z',
|
||||||
|
availability: 'Public',
|
||||||
|
images: [{ url: 'https://image.civitai.com/preview.jpg' }],
|
||||||
|
files: [{ id: 1, type: 'Model', name: 'model.safetensors', sizeKB: 1000 }],
|
||||||
|
modelSizeKB: 1000,
|
||||||
|
existsLocally: false,
|
||||||
|
hasBeenDownloaded: false,
|
||||||
|
};
|
||||||
|
}
|
||||||
|
|
||||||
|
// Newest-first, matching the Civitai API response order. The default selection
|
||||||
|
// takes the first version, so the newest (largest id) must come first.
|
||||||
|
const versions = [makeVersion(250, 'v3'), makeVersion(100, 'v1'), makeVersion(30, 'v0')];
|
||||||
|
|
||||||
|
describe('DownloadManager latest-version default', () => {
|
||||||
|
let DownloadManager;
|
||||||
|
let manager;
|
||||||
|
|
||||||
|
beforeEach(async () => {
|
||||||
|
document.body.innerHTML = '';
|
||||||
|
setupDownloadDom();
|
||||||
|
|
||||||
|
mockApiClient.fetchCivitaiVersions.mockReset();
|
||||||
|
mockLoadingManager.showSimpleLoading.mockClear();
|
||||||
|
mockLoadingManager.hide.mockClear();
|
||||||
|
|
||||||
|
vi.resetModules();
|
||||||
|
({ DownloadManager } = await import(DOWNLOAD_MANAGER_MODULE));
|
||||||
|
manager = new DownloadManager();
|
||||||
|
manager.apiClient = mockApiClient;
|
||||||
|
});
|
||||||
|
|
||||||
|
afterEach(() => {
|
||||||
|
document.body.innerHTML = '';
|
||||||
|
});
|
||||||
|
|
||||||
|
describe('single URL without modelVersionId', () => {
|
||||||
|
it('auto-selects the latest version so no manual version click is needed', async () => {
|
||||||
|
document.getElementById('modelUrl').value = 'https://civitai.red/models/837884/midjourney-artful-nsfw';
|
||||||
|
mockApiClient.fetchCivitaiVersions.mockResolvedValue(versions);
|
||||||
|
|
||||||
|
await manager.validateAndFetchVersions();
|
||||||
|
|
||||||
|
expect(manager.modelId).toBe('837884');
|
||||||
|
expect(manager.currentVersion.id).toBe(250);
|
||||||
|
// The version step is shown with the latest version pre-selected.
|
||||||
|
expect(document.getElementById('versionStep').style.display).toBe('block');
|
||||||
|
const selected = document.querySelector('.version-item.selected');
|
||||||
|
expect(selected.dataset.versionId).toBe('250');
|
||||||
|
expect(document.getElementById('nextFromVersion').disabled).toBe(false);
|
||||||
|
});
|
||||||
|
|
||||||
|
it('still honours an explicit modelVersionId from the URL', async () => {
|
||||||
|
document.getElementById('modelUrl').value =
|
||||||
|
'https://civitai.red/models/837884/midjourney-artful-nsfw?modelVersionId=30';
|
||||||
|
mockApiClient.fetchCivitaiVersions.mockResolvedValue(versions);
|
||||||
|
|
||||||
|
await manager.validateAndFetchVersions();
|
||||||
|
|
||||||
|
expect(manager.currentVersion.id).toBe(30);
|
||||||
|
});
|
||||||
|
});
|
||||||
|
|
||||||
|
describe('fetchVersionsForCurrentModel without modelVersionId', () => {
|
||||||
|
it('auto-selects the latest version', async () => {
|
||||||
|
manager.modelId = '837884';
|
||||||
|
manager.modelVersionId = null;
|
||||||
|
mockApiClient.fetchCivitaiVersions.mockResolvedValue(versions);
|
||||||
|
|
||||||
|
await manager.fetchVersionsForCurrentModel();
|
||||||
|
|
||||||
|
expect(manager.currentVersion.id).toBe(250);
|
||||||
|
});
|
||||||
|
|
||||||
|
it('still honours an explicit modelVersionId', async () => {
|
||||||
|
manager.modelId = '837884';
|
||||||
|
manager.modelVersionId = '100';
|
||||||
|
mockApiClient.fetchCivitaiVersions.mockResolvedValue(versions);
|
||||||
|
|
||||||
|
await manager.fetchVersionsForCurrentModel();
|
||||||
|
|
||||||
|
expect(manager.currentVersion.id).toBe(100);
|
||||||
|
});
|
||||||
|
});
|
||||||
|
|
||||||
|
describe('multi-URL batch without modelVersionId', () => {
|
||||||
|
it('defaults each item to its latest version (first in API response)', async () => {
|
||||||
|
document.getElementById('modelUrl').value =
|
||||||
|
'https://civitai.red/models/111/foo\nhttps://civitai.red/models/222/bar';
|
||||||
|
const versionsA = [makeVersion(40, 'a2'), makeVersion(10, 'a1')];
|
||||||
|
const versionsB = [makeVersion(500, 'b1'), makeVersion(7, 'b2')];
|
||||||
|
mockApiClient.fetchCivitaiVersions.mockImplementation(async modelId =>
|
||||||
|
modelId === '111' ? versionsA : versionsB
|
||||||
|
);
|
||||||
|
|
||||||
|
await manager.validateAndFetchVersions();
|
||||||
|
|
||||||
|
expect(manager.isBatchMode).toBe(true);
|
||||||
|
expect(manager.batchModels).toHaveLength(2);
|
||||||
|
expect(manager.batchModels[0].selectedVersion.id).toBe(40);
|
||||||
|
expect(manager.batchModels[1].selectedVersion.id).toBe(500);
|
||||||
|
});
|
||||||
|
|
||||||
|
it('still honours an explicit modelVersionId per URL', async () => {
|
||||||
|
document.getElementById('modelUrl').value =
|
||||||
|
'https://civitai.red/models/111/foo?modelVersionId=10\nhttps://civitai.red/models/222/bar';
|
||||||
|
const versionsA = [makeVersion(40, 'a2'), makeVersion(10, 'a1')];
|
||||||
|
const versionsB = [makeVersion(500, 'b1'), makeVersion(7, 'b2')];
|
||||||
|
mockApiClient.fetchCivitaiVersions.mockImplementation(async modelId =>
|
||||||
|
modelId === '111' ? versionsA : versionsB
|
||||||
|
);
|
||||||
|
|
||||||
|
await manager.validateAndFetchVersions();
|
||||||
|
|
||||||
|
expect(manager.batchModels[0].selectedVersion.id).toBe(10);
|
||||||
|
expect(manager.batchModels[1].selectedVersion.id).toBe(500);
|
||||||
|
});
|
||||||
|
});
|
||||||
|
});
|
||||||
@@ -501,3 +501,33 @@ describe('SettingsManager library controls', () => {
|
|||||||
expect(document.getElementById('exampleImagesUriTemplateSetting').style.display).toBe('none');
|
expect(document.getElementById('exampleImagesUriTemplateSetting').style.display).toBe('none');
|
||||||
});
|
});
|
||||||
});
|
});
|
||||||
|
|
||||||
|
describe('SettingsManager recipes layout switch', () => {
|
||||||
|
it('dispatches lm:recipes-layout-changed without recalculating the old scroller', async () => {
|
||||||
|
const manager = createManager();
|
||||||
|
const select = document.createElement('select');
|
||||||
|
select.id = 'recipesLayout';
|
||||||
|
const option = document.createElement('option');
|
||||||
|
option.value = 'masonry';
|
||||||
|
select.appendChild(option);
|
||||||
|
select.value = 'masonry';
|
||||||
|
document.body.appendChild(select);
|
||||||
|
|
||||||
|
const calculateLayout = vi.fn();
|
||||||
|
state.virtualScroller = { calculateLayout };
|
||||||
|
|
||||||
|
const dispatchSpy = vi.spyOn(window, 'dispatchEvent');
|
||||||
|
|
||||||
|
await manager.saveSelectSetting('recipesLayout', 'recipes_layout');
|
||||||
|
|
||||||
|
const layoutEvent = dispatchSpy.mock.calls
|
||||||
|
.map(([event]) => event)
|
||||||
|
.find(event => event.type === 'lm:recipes-layout-changed');
|
||||||
|
expect(layoutEvent).toBeInstanceOf(CustomEvent);
|
||||||
|
expect(calculateLayout).not.toHaveBeenCalled();
|
||||||
|
expect(showToast).not.toHaveBeenCalled();
|
||||||
|
|
||||||
|
dispatchSpy.mockRestore();
|
||||||
|
delete state.virtualScroller;
|
||||||
|
});
|
||||||
|
});
|
||||||
|
|||||||
@@ -69,6 +69,7 @@ vi.mock('../../../static/js/components/DuplicatesManager.js', () => ({
|
|||||||
|
|
||||||
vi.mock('../../../static/js/utils/infiniteScroll.js', () => ({
|
vi.mock('../../../static/js/utils/infiniteScroll.js', () => ({
|
||||||
refreshVirtualScroll: refreshVirtualScrollMock,
|
refreshVirtualScroll: refreshVirtualScrollMock,
|
||||||
|
recreateVirtualScroll: vi.fn(),
|
||||||
}));
|
}));
|
||||||
|
|
||||||
vi.mock('../../../static/js/api/recipeApi.js', () => ({
|
vi.mock('../../../static/js/api/recipeApi.js', () => ({
|
||||||
|
|||||||
@@ -17,7 +17,8 @@ describe('state module', () => {
|
|||||||
civitai_host: 'civitai.com',
|
civitai_host: 'civitai.com',
|
||||||
language: 'en',
|
language: 'en',
|
||||||
blur_mature_content: true,
|
blur_mature_content: true,
|
||||||
mature_blur_level: 'R'
|
mature_blur_level: 'R',
|
||||||
|
recipes_layout: 'grid'
|
||||||
});
|
});
|
||||||
|
|
||||||
expect(defaultSettings.download_path_templates).toEqual(DEFAULT_PATH_TEMPLATES);
|
expect(defaultSettings.download_path_templates).toEqual(DEFAULT_PATH_TEMPLATES);
|
||||||
|
|||||||
@@ -0,0 +1,198 @@
|
|||||||
|
import { afterEach, beforeEach, describe, expect, it, vi } from 'vitest';
|
||||||
|
|
||||||
|
const {
|
||||||
|
VirtualScrollerMock,
|
||||||
|
MasonryScrollerMock,
|
||||||
|
RecipeCardMock,
|
||||||
|
fetchRecipesPageMock,
|
||||||
|
getModelApiClientMock,
|
||||||
|
} = vi.hoisted(() => {
|
||||||
|
const makeScrollerClass = () => vi.fn(function (options) {
|
||||||
|
this.options = options;
|
||||||
|
this.initialize = vi.fn(async () => {});
|
||||||
|
this.dispose = vi.fn();
|
||||||
|
this.reset = vi.fn();
|
||||||
|
this.handlePageUpDown = vi.fn();
|
||||||
|
});
|
||||||
|
|
||||||
|
return {
|
||||||
|
VirtualScrollerMock: makeScrollerClass(),
|
||||||
|
MasonryScrollerMock: makeScrollerClass(),
|
||||||
|
RecipeCardMock: vi.fn(function () {
|
||||||
|
this.element = document.createElement('div');
|
||||||
|
}),
|
||||||
|
fetchRecipesPageMock: vi.fn(async () => ({ items: [], totalItems: 0, hasMore: false })),
|
||||||
|
getModelApiClientMock: vi.fn(() => ({
|
||||||
|
fetchModelsPage: vi.fn(async () => ({ items: [], totalItems: 0, hasMore: false })),
|
||||||
|
})),
|
||||||
|
};
|
||||||
|
});
|
||||||
|
|
||||||
|
vi.mock('../../../static/js/utils/VirtualScroller.js', () => ({
|
||||||
|
VirtualScroller: VirtualScrollerMock,
|
||||||
|
}));
|
||||||
|
|
||||||
|
vi.mock('../../../static/js/utils/MasonryScroller.js', () => ({
|
||||||
|
MasonryScroller: MasonryScrollerMock,
|
||||||
|
}));
|
||||||
|
|
||||||
|
vi.mock('../../../static/js/components/RecipeCard.js', () => ({
|
||||||
|
RecipeCard: RecipeCardMock,
|
||||||
|
}));
|
||||||
|
|
||||||
|
vi.mock('../../../static/js/api/recipeApi.js', () => ({
|
||||||
|
fetchRecipesPage: fetchRecipesPageMock,
|
||||||
|
}));
|
||||||
|
|
||||||
|
vi.mock('../../../static/js/api/modelApiFactory.js', () => ({
|
||||||
|
getModelApiClient: getModelApiClientMock,
|
||||||
|
}));
|
||||||
|
|
||||||
|
vi.mock('../../../static/js/components/shared/ModelCard.js', () => ({
|
||||||
|
createModelCard: vi.fn(() => document.createElement('div')),
|
||||||
|
setupModelCardEventDelegation: vi.fn(),
|
||||||
|
}));
|
||||||
|
|
||||||
|
vi.mock('../../../static/js/utils/uiHelpers.js', () => ({
|
||||||
|
showToast: vi.fn(),
|
||||||
|
}));
|
||||||
|
|
||||||
|
import {
|
||||||
|
initializeInfiniteScroll,
|
||||||
|
recreateVirtualScroll,
|
||||||
|
} from '../../../static/js/utils/infiniteScroll.js';
|
||||||
|
import { state } from '../../../static/js/state/index.js';
|
||||||
|
|
||||||
|
function setupPageDom() {
|
||||||
|
const pageContent = document.createElement('div');
|
||||||
|
pageContent.className = 'page-content';
|
||||||
|
const container = document.createElement('div');
|
||||||
|
container.className = 'container';
|
||||||
|
pageContent.appendChild(container);
|
||||||
|
document.body.appendChild(pageContent);
|
||||||
|
|
||||||
|
const grid = document.createElement('div');
|
||||||
|
vi.spyOn(document, 'getElementById').mockImplementation((id) =>
|
||||||
|
id === 'recipeGrid' || id === 'modelGrid' ? grid : null);
|
||||||
|
|
||||||
|
return { grid };
|
||||||
|
}
|
||||||
|
|
||||||
|
describe('infiniteScroll scroller class branching', () => {
|
||||||
|
let originalSettings;
|
||||||
|
|
||||||
|
beforeEach(() => {
|
||||||
|
originalSettings = state.global.settings;
|
||||||
|
state.global.settings = { ...originalSettings };
|
||||||
|
if (state.pages.recipes) {
|
||||||
|
state.pages.recipes.duplicatesMode = false;
|
||||||
|
}
|
||||||
|
state.virtualScroller = null;
|
||||||
|
state.keyboardNavHandler = null;
|
||||||
|
setupPageDom();
|
||||||
|
});
|
||||||
|
|
||||||
|
afterEach(() => {
|
||||||
|
state.global.settings = originalSettings;
|
||||||
|
state.virtualScroller = null;
|
||||||
|
state.keyboardNavHandler = null;
|
||||||
|
document.body.innerHTML = '';
|
||||||
|
vi.restoreAllMocks();
|
||||||
|
vi.clearAllMocks();
|
||||||
|
});
|
||||||
|
|
||||||
|
it('constructs MasonryScroller for the recipes page when recipes_layout is masonry', async () => {
|
||||||
|
state.global.settings.recipes_layout = 'masonry';
|
||||||
|
|
||||||
|
await initializeInfiniteScroll('recipes');
|
||||||
|
|
||||||
|
expect(MasonryScrollerMock).toHaveBeenCalledTimes(1);
|
||||||
|
expect(VirtualScrollerMock).not.toHaveBeenCalled();
|
||||||
|
expect(state.virtualScroller).toBeInstanceOf(MasonryScrollerMock);
|
||||||
|
});
|
||||||
|
|
||||||
|
it('constructs VirtualScroller for the recipes page when recipes_layout is grid', async () => {
|
||||||
|
state.global.settings.recipes_layout = 'grid';
|
||||||
|
|
||||||
|
await initializeInfiniteScroll('recipes');
|
||||||
|
|
||||||
|
expect(VirtualScrollerMock).toHaveBeenCalledTimes(1);
|
||||||
|
expect(MasonryScrollerMock).not.toHaveBeenCalled();
|
||||||
|
expect(state.virtualScroller).toBeInstanceOf(VirtualScrollerMock);
|
||||||
|
});
|
||||||
|
|
||||||
|
it('falls back to the grid branch when recipes_layout is missing', async () => {
|
||||||
|
delete state.global.settings.recipes_layout;
|
||||||
|
|
||||||
|
await expect(initializeInfiniteScroll('recipes')).resolves.toBeUndefined();
|
||||||
|
|
||||||
|
expect(VirtualScrollerMock).toHaveBeenCalledTimes(1);
|
||||||
|
expect(MasonryScrollerMock).not.toHaveBeenCalled();
|
||||||
|
});
|
||||||
|
|
||||||
|
it('always constructs VirtualScroller for the loras page regardless of recipes_layout', async () => {
|
||||||
|
state.global.settings.recipes_layout = 'masonry';
|
||||||
|
|
||||||
|
await initializeInfiniteScroll('loras');
|
||||||
|
|
||||||
|
expect(VirtualScrollerMock).toHaveBeenCalledTimes(1);
|
||||||
|
expect(MasonryScrollerMock).not.toHaveBeenCalled();
|
||||||
|
});
|
||||||
|
});
|
||||||
|
|
||||||
|
describe('recreateVirtualScroll', () => {
|
||||||
|
let originalSettings;
|
||||||
|
|
||||||
|
beforeEach(() => {
|
||||||
|
originalSettings = state.global.settings;
|
||||||
|
state.global.settings = { ...originalSettings };
|
||||||
|
if (state.pages.recipes) {
|
||||||
|
state.pages.recipes.duplicatesMode = false;
|
||||||
|
}
|
||||||
|
state.virtualScroller = null;
|
||||||
|
state.keyboardNavHandler = null;
|
||||||
|
setupPageDom();
|
||||||
|
});
|
||||||
|
|
||||||
|
afterEach(() => {
|
||||||
|
state.global.settings = originalSettings;
|
||||||
|
state.virtualScroller = null;
|
||||||
|
state.keyboardNavHandler = null;
|
||||||
|
document.body.innerHTML = '';
|
||||||
|
vi.restoreAllMocks();
|
||||||
|
vi.clearAllMocks();
|
||||||
|
});
|
||||||
|
|
||||||
|
it('cleans up keyboard navigation, disposes the old scroller, and rebuilds with the new layout', async () => {
|
||||||
|
state.global.settings.recipes_layout = 'grid';
|
||||||
|
await initializeInfiniteScroll('recipes');
|
||||||
|
const oldScroller = state.virtualScroller;
|
||||||
|
const oldKeyboardHandler = state.keyboardNavHandler;
|
||||||
|
expect(oldKeyboardHandler).toBeTypeOf('function');
|
||||||
|
|
||||||
|
const removeEventListenerSpy = vi.spyOn(document, 'removeEventListener');
|
||||||
|
state.global.settings.recipes_layout = 'masonry';
|
||||||
|
|
||||||
|
await recreateVirtualScroll('recipes');
|
||||||
|
|
||||||
|
// cleanupKeyboardNavigation removed the previous document keydown listener
|
||||||
|
expect(removeEventListenerSpy).toHaveBeenCalledWith('keydown', oldKeyboardHandler);
|
||||||
|
// Old scroller disposed and replaced by a new masonry instance
|
||||||
|
expect(oldScroller.dispose).toHaveBeenCalledTimes(1);
|
||||||
|
expect(MasonryScrollerMock).toHaveBeenCalledTimes(1);
|
||||||
|
expect(state.virtualScroller).toBeInstanceOf(MasonryScrollerMock);
|
||||||
|
expect(state.virtualScroller).not.toBe(oldScroller);
|
||||||
|
// A fresh keyboard navigation listener was registered for the new instance
|
||||||
|
expect(state.keyboardNavHandler).toBeTypeOf('function');
|
||||||
|
expect(state.keyboardNavHandler).not.toBe(oldKeyboardHandler);
|
||||||
|
});
|
||||||
|
|
||||||
|
it('works when there is no existing virtual scroller', async () => {
|
||||||
|
state.global.settings.recipes_layout = 'masonry';
|
||||||
|
|
||||||
|
await expect(recreateVirtualScroll('recipes')).resolves.toBeUndefined();
|
||||||
|
|
||||||
|
expect(MasonryScrollerMock).toHaveBeenCalledTimes(1);
|
||||||
|
expect(state.virtualScroller).toBeInstanceOf(MasonryScrollerMock);
|
||||||
|
});
|
||||||
|
});
|
||||||
@@ -0,0 +1,627 @@
|
|||||||
|
import { afterEach, beforeEach, describe, expect, it, vi } from 'vitest';
|
||||||
|
|
||||||
|
import { MasonryScroller } from '../../../static/js/utils/MasonryScroller.js';
|
||||||
|
import { VirtualScroller } from '../../../static/js/utils/VirtualScroller.js';
|
||||||
|
import { getCurrentPageState, setCurrentPageType } from '../../../static/js/state/index.js';
|
||||||
|
|
||||||
|
// jsdom does not always provide requestAnimationFrame; polyfill when missing
|
||||||
|
if (typeof window !== 'undefined' && typeof window.requestAnimationFrame !== 'function') {
|
||||||
|
window.requestAnimationFrame = (cb) => setTimeout(cb, 0);
|
||||||
|
window.cancelAnimationFrame = (id) => clearTimeout(id);
|
||||||
|
}
|
||||||
|
|
||||||
|
const CONTAINER_WIDTH = 768; // yields 3 columns at default density: floor((768+12)/(240+12)) = 3
|
||||||
|
const COLUMN_GAP = 12;
|
||||||
|
const ROW_GAP = 20;
|
||||||
|
const PAD_TOP = 4;
|
||||||
|
const PAD_BOTTOM = 4;
|
||||||
|
const ITEM_WIDTH = (CONTAINER_WIDTH - 2 * COLUMN_GAP) / 3; // 248
|
||||||
|
const FALLBACK_HEIGHT = ITEM_WIDTH / (896 / 1152);
|
||||||
|
|
||||||
|
function createItemFn() {
|
||||||
|
const el = document.createElement('div');
|
||||||
|
const card = document.createElement('div');
|
||||||
|
card.className = 'model-card';
|
||||||
|
const preview = document.createElement('div');
|
||||||
|
preview.className = 'card-preview';
|
||||||
|
card.appendChild(preview);
|
||||||
|
el.appendChild(card);
|
||||||
|
return el;
|
||||||
|
}
|
||||||
|
|
||||||
|
function makeItems(dimensions) {
|
||||||
|
return dimensions.map((dims, i) => ({
|
||||||
|
file_path: `/recipes/item-${i}.png`,
|
||||||
|
...dims,
|
||||||
|
}));
|
||||||
|
}
|
||||||
|
|
||||||
|
/**
|
||||||
|
* Build a scroller attached to a stubbed container. clientWidth/clientHeight
|
||||||
|
* are 0 in jsdom, so they are defined explicitly for deterministic layout.
|
||||||
|
*/
|
||||||
|
function createScroller({ items = [], fetchItemsFn, overscan, viewportHeight = 600, createItemFn: customCreateItemFn } = {}) {
|
||||||
|
const wrapper = document.createElement('div');
|
||||||
|
Object.defineProperty(wrapper, 'clientWidth', { value: CONTAINER_WIDTH, configurable: true });
|
||||||
|
Object.defineProperty(wrapper, 'clientHeight', { value: viewportHeight, configurable: true });
|
||||||
|
|
||||||
|
const grid = document.createElement('div');
|
||||||
|
wrapper.appendChild(grid);
|
||||||
|
document.body.appendChild(wrapper);
|
||||||
|
|
||||||
|
const fetchMock = fetchItemsFn || vi.fn(async () => ({ items, totalItems: items.length, hasMore: false }));
|
||||||
|
|
||||||
|
const scroller = new MasonryScroller({
|
||||||
|
gridElement: grid,
|
||||||
|
containerElement: wrapper,
|
||||||
|
scrollContainer: wrapper,
|
||||||
|
createItemFn: customCreateItemFn || createItemFn,
|
||||||
|
fetchItemsFn: fetchMock,
|
||||||
|
overscan,
|
||||||
|
});
|
||||||
|
|
||||||
|
return { scroller, wrapper, grid, fetchMock };
|
||||||
|
}
|
||||||
|
|
||||||
|
describe('MasonryScroller', () => {
|
||||||
|
const liveScrollers = [];
|
||||||
|
|
||||||
|
beforeEach(() => {
|
||||||
|
setCurrentPageType('recipes');
|
||||||
|
getCurrentPageState().duplicatesMode = false;
|
||||||
|
});
|
||||||
|
|
||||||
|
afterEach(() => {
|
||||||
|
while (liveScrollers.length > 0) {
|
||||||
|
liveScrollers.pop().dispose();
|
||||||
|
}
|
||||||
|
getCurrentPageState().duplicatesMode = false;
|
||||||
|
});
|
||||||
|
|
||||||
|
function track(setup) {
|
||||||
|
liveScrollers.push(setup.scroller);
|
||||||
|
return setup;
|
||||||
|
}
|
||||||
|
|
||||||
|
it('adds virtual-scroll and masonry-layout classes and creates the spacer', () => {
|
||||||
|
const { scroller, grid } = track(createScroller());
|
||||||
|
|
||||||
|
expect(grid.classList.contains('virtual-scroll')).toBe(true);
|
||||||
|
expect(grid.classList.contains('masonry-layout')).toBe(true);
|
||||||
|
expect(scroller.spacerElement.className).toBe('virtual-scroll-spacer');
|
||||||
|
expect(grid.style.position).toBe('relative');
|
||||||
|
expect(grid.contains(scroller.spacerElement)).toBe(true);
|
||||||
|
});
|
||||||
|
|
||||||
|
it('computes density-based column count and item width', () => {
|
||||||
|
const { scroller } = track(createScroller());
|
||||||
|
|
||||||
|
expect(scroller.columnsCount).toBe(3);
|
||||||
|
expect(scroller.itemWidth).toBeCloseTo(ITEM_WIDTH);
|
||||||
|
});
|
||||||
|
|
||||||
|
it('places each item into the shortest column', () => {
|
||||||
|
// Heights at ITEM_WIDTH=248: 496, 248, 124, 248, 248, 248
|
||||||
|
const items = makeItems([
|
||||||
|
{ width: 100, height: 200 },
|
||||||
|
{ width: 100, height: 100 },
|
||||||
|
{ width: 100, height: 50 },
|
||||||
|
{ width: 100, height: 100 },
|
||||||
|
{ width: 100, height: 100 },
|
||||||
|
{ width: 100, height: 100 },
|
||||||
|
]);
|
||||||
|
const { scroller } = track(createScroller({ items }));
|
||||||
|
|
||||||
|
scroller.refreshWithData(items, items.length, false);
|
||||||
|
|
||||||
|
const cols = scroller.positions.map((p) => p.col);
|
||||||
|
expect(cols).toEqual([0, 1, 2, 2, 1, 2]);
|
||||||
|
|
||||||
|
// Tops follow the accumulated shortest-column heights
|
||||||
|
expect(scroller.positions[0].top).toBeCloseTo(PAD_TOP);
|
||||||
|
expect(scroller.positions[1].top).toBeCloseTo(PAD_TOP);
|
||||||
|
expect(scroller.positions[2].top).toBeCloseTo(PAD_TOP);
|
||||||
|
expect(scroller.positions[3].top).toBeCloseTo(PAD_TOP + 124 + ROW_GAP); // 148
|
||||||
|
expect(scroller.positions[4].top).toBeCloseTo(PAD_TOP + 248 + ROW_GAP); // 272
|
||||||
|
expect(scroller.positions[5].top).toBeCloseTo(148 + 248 + ROW_GAP); // 416
|
||||||
|
|
||||||
|
// Left offsets are column index * (itemWidth + columnGap)
|
||||||
|
expect(scroller.positions[1].left).toBeCloseTo(ITEM_WIDTH + COLUMN_GAP);
|
||||||
|
expect(scroller.positions[2].left).toBeCloseTo(2 * (ITEM_WIDTH + COLUMN_GAP));
|
||||||
|
});
|
||||||
|
|
||||||
|
it('sets spacer height to the tallest column minus trailing gap plus padding', () => {
|
||||||
|
const items = makeItems([
|
||||||
|
{ width: 100, height: 200 },
|
||||||
|
{ width: 100, height: 100 },
|
||||||
|
{ width: 100, height: 50 },
|
||||||
|
{ width: 100, height: 100 },
|
||||||
|
{ width: 100, height: 100 },
|
||||||
|
{ width: 100, height: 100 },
|
||||||
|
]);
|
||||||
|
const { scroller } = track(createScroller({ items }));
|
||||||
|
|
||||||
|
scroller.refreshWithData(items, items.length, false);
|
||||||
|
|
||||||
|
// Column heights after placement: [520, 540, 684] (each starts at padTop)
|
||||||
|
const expected = 684 - ROW_GAP + PAD_TOP + PAD_BOTTOM; // 672
|
||||||
|
expect(scroller.spacerElement.style.height).toBe(`${expected}px`);
|
||||||
|
});
|
||||||
|
|
||||||
|
it('falls back to the 896/1152 ratio for items without dimensions', () => {
|
||||||
|
const items = makeItems([{}, { width: 100, height: 100 }]);
|
||||||
|
const { scroller } = track(createScroller({ items }));
|
||||||
|
|
||||||
|
scroller.refreshWithData(items, items.length, false);
|
||||||
|
|
||||||
|
expect(scroller.positions[0].height).toBeCloseTo(FALLBACK_HEIGHT);
|
||||||
|
expect(scroller.positions[1].height).toBeCloseTo(ITEM_WIDTH);
|
||||||
|
});
|
||||||
|
|
||||||
|
it('visible range respects the overscan band', () => {
|
||||||
|
const items = makeItems([
|
||||||
|
{ width: 100, height: 200 },
|
||||||
|
{ width: 100, height: 100 },
|
||||||
|
{ width: 100, height: 50 },
|
||||||
|
{ width: 100, height: 100 },
|
||||||
|
{ width: 100, height: 100 },
|
||||||
|
{ width: 100, height: 100 },
|
||||||
|
]);
|
||||||
|
const { scroller, wrapper } = track(createScroller({ items, viewportHeight: 300 }));
|
||||||
|
|
||||||
|
scroller.refreshWithData(items, items.length, false);
|
||||||
|
wrapper.scrollTop = 0;
|
||||||
|
|
||||||
|
// Without overscan, item 5 (top=416) is below the 300px viewport
|
||||||
|
scroller.overscan = 0;
|
||||||
|
let visible = scroller.getVisibleRange();
|
||||||
|
expect([...visible].sort((a, b) => a - b)).toEqual([0, 1, 2, 3, 4]);
|
||||||
|
|
||||||
|
// overscan=1 extends the band by one itemWidth (248px), pulling item 5 in
|
||||||
|
scroller.overscan = 1;
|
||||||
|
visible = scroller.getVisibleRange();
|
||||||
|
expect(visible.has(5)).toBe(true);
|
||||||
|
});
|
||||||
|
|
||||||
|
it('renders visible items with inline masonry styles and clears model-card max-width', () => {
|
||||||
|
const items = makeItems([
|
||||||
|
{ width: 100, height: 200 },
|
||||||
|
{ width: 100, height: 100 },
|
||||||
|
{ width: 100, height: 50 },
|
||||||
|
]);
|
||||||
|
const { scroller, grid, wrapper } = track(createScroller({ items, viewportHeight: 3000 }));
|
||||||
|
|
||||||
|
scroller.refreshWithData(items, items.length, false);
|
||||||
|
wrapper.scrollTop = 0;
|
||||||
|
scroller.overscan = 5;
|
||||||
|
|
||||||
|
scroller.renderItems();
|
||||||
|
|
||||||
|
const rendered = grid.querySelectorAll('.virtual-scroll-item');
|
||||||
|
expect(rendered.length).toBe(3);
|
||||||
|
rendered.forEach((el, i) => {
|
||||||
|
expect(el.style.position).toBe('absolute');
|
||||||
|
expect(el.style.width).toBe(`${scroller.positions[i].width}px`);
|
||||||
|
expect(el.style.height).toBe(`${scroller.positions[i].height}px`);
|
||||||
|
expect(el.style.left).toBe(`${scroller.positions[i].left}px`);
|
||||||
|
expect(el.style.top).toBe(`${scroller.positions[i].top}px`);
|
||||||
|
expect(el.querySelector('.model-card').style.maxWidth).toBe('none');
|
||||||
|
expect(el.querySelector('.model-card').style.minWidth).toBe('0');
|
||||||
|
});
|
||||||
|
});
|
||||||
|
|
||||||
|
it('clears model-card max-width/min-width when the item element is the card root', () => {
|
||||||
|
// Production shape for recipe cards: RecipeCard returns the .model-card
|
||||||
|
// element itself, so the scroller must clear constraints on the element
|
||||||
|
// rather than a descendant (querySelector would find nothing).
|
||||||
|
const items = makeItems([{ width: 100, height: 200 }]);
|
||||||
|
const cardRoot = document.createElement('div');
|
||||||
|
cardRoot.className = 'model-card';
|
||||||
|
const { scroller, grid, wrapper } = track(createScroller({
|
||||||
|
items,
|
||||||
|
viewportHeight: 3000,
|
||||||
|
createItemFn: () => cardRoot.cloneNode(true),
|
||||||
|
}));
|
||||||
|
|
||||||
|
scroller.refreshWithData(items, items.length, false);
|
||||||
|
wrapper.scrollTop = 0;
|
||||||
|
scroller.overscan = 5;
|
||||||
|
|
||||||
|
scroller.renderItems();
|
||||||
|
|
||||||
|
const rendered = grid.querySelectorAll('.virtual-scroll-item');
|
||||||
|
expect(rendered.length).toBe(1);
|
||||||
|
expect(rendered[0].style.maxWidth).toBe('none');
|
||||||
|
expect(rendered[0].style.minWidth).toBe('0');
|
||||||
|
});
|
||||||
|
|
||||||
|
it('triggers loadMoreItems when scrolled to the bottom', async () => {
|
||||||
|
const firstPage = makeItems([
|
||||||
|
{ width: 100, height: 200 },
|
||||||
|
{ width: 100, height: 100 },
|
||||||
|
{ width: 100, height: 50 },
|
||||||
|
{ width: 100, height: 100 },
|
||||||
|
{ width: 100, height: 100 },
|
||||||
|
{ width: 100, height: 100 },
|
||||||
|
]);
|
||||||
|
const fetchMock = vi.fn(async () => ({ items: makeItems([{ width: 100, height: 100 }]), totalItems: 7, hasMore: false }));
|
||||||
|
const { scroller, wrapper } = track(createScroller({ fetchItemsFn: fetchMock, viewportHeight: 600 }));
|
||||||
|
|
||||||
|
scroller.refreshWithData(firstPage, 100, true);
|
||||||
|
const pageState = getCurrentPageState();
|
||||||
|
const expectedPage = pageState.currentPage;
|
||||||
|
|
||||||
|
// contentBottom = 664; scrollBottom = 100 + 600 = 700 >= 664 - threshold
|
||||||
|
wrapper.scrollTop = 100;
|
||||||
|
scroller.handleScroll();
|
||||||
|
|
||||||
|
await vi.waitFor(() => {
|
||||||
|
expect(fetchMock).toHaveBeenCalledWith(expectedPage, scroller.pageSize);
|
||||||
|
});
|
||||||
|
});
|
||||||
|
|
||||||
|
it('does not trigger loadMoreItems when far from the bottom', async () => {
|
||||||
|
const fetchMock = vi.fn(async () => ({ items: [], totalItems: 0, hasMore: false }));
|
||||||
|
const manyItems = makeItems(Array.from({ length: 30 }, () => ({ width: 100, height: 200 })));
|
||||||
|
const { scroller, wrapper } = track(createScroller({ fetchItemsFn: fetchMock, viewportHeight: 600 }));
|
||||||
|
|
||||||
|
scroller.refreshWithData(manyItems, 1000, true);
|
||||||
|
wrapper.scrollTop = 0;
|
||||||
|
scroller.handleScroll();
|
||||||
|
|
||||||
|
await new Promise((resolve) => setTimeout(resolve, 20));
|
||||||
|
expect(fetchMock).not.toHaveBeenCalled();
|
||||||
|
});
|
||||||
|
|
||||||
|
it('returns false from calculateLayout in duplicates mode', () => {
|
||||||
|
const { scroller } = track(createScroller());
|
||||||
|
|
||||||
|
getCurrentPageState().duplicatesMode = true;
|
||||||
|
expect(scroller.calculateLayout()).toBe(false);
|
||||||
|
});
|
||||||
|
|
||||||
|
it('computes placement and spacer synchronously during initialize (before rAF)', async () => {
|
||||||
|
const items = makeItems([
|
||||||
|
{ width: 100, height: 200 },
|
||||||
|
{ width: 100, height: 100 },
|
||||||
|
{ width: 100, height: 50 },
|
||||||
|
{ width: 100, height: 100 },
|
||||||
|
{ width: 100, height: 100 },
|
||||||
|
{ width: 100, height: 100 },
|
||||||
|
]);
|
||||||
|
const { scroller } = track(createScroller({ items }));
|
||||||
|
|
||||||
|
await scroller.initialize();
|
||||||
|
|
||||||
|
// Assert immediately after initialize resolves, without waiting for rAF
|
||||||
|
expect(scroller.positions.length).toBe(items.length);
|
||||||
|
const expected = 684 - ROW_GAP + PAD_TOP + PAD_BOTTOM; // 672
|
||||||
|
expect(scroller.spacerElement.style.height).toBe(`${expected}px`);
|
||||||
|
expect(getCurrentPageState().currentPage).toBe(2);
|
||||||
|
});
|
||||||
|
|
||||||
|
it('shows the error placeholder and resets isLoading when the initial fetch fails', async () => {
|
||||||
|
const fetchMock = vi.fn(async () => {
|
||||||
|
throw new Error('network down');
|
||||||
|
});
|
||||||
|
const { scroller, grid } = track(createScroller({ fetchItemsFn: fetchMock }));
|
||||||
|
|
||||||
|
await expect(scroller.initialize()).resolves.toBeUndefined();
|
||||||
|
|
||||||
|
const placeholder = grid.querySelector('#virtualScrollPlaceholder');
|
||||||
|
expect(placeholder).not.toBeNull();
|
||||||
|
expect(placeholder.textContent).toContain('Failed to load items');
|
||||||
|
expect(scroller.isLoading).toBe(false);
|
||||||
|
});
|
||||||
|
|
||||||
|
it('shows the recipes empty placeholder when no items are returned', async () => {
|
||||||
|
const { scroller, grid } = track(createScroller({ items: [] }));
|
||||||
|
|
||||||
|
await scroller.initialize();
|
||||||
|
|
||||||
|
const placeholder = grid.querySelector('#virtualScrollPlaceholder');
|
||||||
|
expect(placeholder).not.toBeNull();
|
||||||
|
expect(placeholder.textContent).toContain('No recipes found');
|
||||||
|
});
|
||||||
|
|
||||||
|
it('dispose removes classes, spacer and event listeners', () => {
|
||||||
|
const { scroller, grid } = track(createScroller());
|
||||||
|
|
||||||
|
scroller.dispose();
|
||||||
|
|
||||||
|
expect(grid.classList.contains('virtual-scroll')).toBe(false);
|
||||||
|
expect(grid.classList.contains('masonry-layout')).toBe(false);
|
||||||
|
expect(grid.querySelector('.virtual-scroll-spacer')).toBeNull();
|
||||||
|
});
|
||||||
|
|
||||||
|
it('exposes every VirtualScroller prototype method (API parity)', () => {
|
||||||
|
const virtualMethods = Object.getOwnPropertyNames(VirtualScroller.prototype);
|
||||||
|
const masonryMethods = new Set(Object.getOwnPropertyNames(MasonryScroller.prototype));
|
||||||
|
|
||||||
|
const missing = virtualMethods.filter((name) => !masonryMethods.has(name));
|
||||||
|
expect(missing).toEqual([]);
|
||||||
|
});
|
||||||
|
|
||||||
|
it('exposes the VirtualScroller property surface after construction', () => {
|
||||||
|
const { scroller } = track(createScroller());
|
||||||
|
|
||||||
|
const expectedProperties = [
|
||||||
|
'items',
|
||||||
|
'renderedItems',
|
||||||
|
'totalItems',
|
||||||
|
'hasMore',
|
||||||
|
'isLoading',
|
||||||
|
'gridElement',
|
||||||
|
'containerElement',
|
||||||
|
'scrollContainer',
|
||||||
|
'columnsCount',
|
||||||
|
'itemWidth',
|
||||||
|
'disabled',
|
||||||
|
'spacerElement',
|
||||||
|
'pageSize',
|
||||||
|
];
|
||||||
|
|
||||||
|
for (const prop of expectedProperties) {
|
||||||
|
expect(scroller[prop]).not.toBeUndefined();
|
||||||
|
}
|
||||||
|
});
|
||||||
|
|
||||||
|
it('updateSingleItem re-places items and shows the updated indicator on rendered cards', () => {
|
||||||
|
// Heights at ITEM_WIDTH=248: 496, 248, 124, 248, 248, 248
|
||||||
|
// Item 2 (height 124) sits in column 2 with items 3 and 5 stacked below it
|
||||||
|
const items = makeItems([
|
||||||
|
{ width: 100, height: 200 },
|
||||||
|
{ width: 100, height: 100 },
|
||||||
|
{ width: 100, height: 50 },
|
||||||
|
{ width: 100, height: 100 },
|
||||||
|
{ width: 100, height: 100 },
|
||||||
|
{ width: 100, height: 100 },
|
||||||
|
]);
|
||||||
|
const { scroller, grid, wrapper } = track(createScroller({ items, viewportHeight: 3000 }));
|
||||||
|
|
||||||
|
scroller.refreshWithData(items, items.length, false);
|
||||||
|
wrapper.scrollTop = 0;
|
||||||
|
scroller.overscan = 5;
|
||||||
|
scroller.renderItems();
|
||||||
|
|
||||||
|
const spacerBefore = scroller.spacerElement.style.height;
|
||||||
|
const colsBefore = scroller.positions.map((p) => p.col);
|
||||||
|
|
||||||
|
const result = scroller.updateSingleItem('/recipes/item-2.png', { width: 100, height: 150 });
|
||||||
|
|
||||||
|
expect(result).toBe(true);
|
||||||
|
|
||||||
|
// Item 2 height grew 124 -> 372, so the full synchronous re-placement
|
||||||
|
// re-flows every later item (item 3 moves from column 2 to column 1)
|
||||||
|
expect(scroller.positions[2].height).toBeCloseTo(ITEM_WIDTH * 1.5);
|
||||||
|
expect(scroller.positions.map((p) => p.col)).not.toEqual(colsBefore);
|
||||||
|
expect(scroller.spacerElement.style.height).not.toBe(spacerBefore);
|
||||||
|
|
||||||
|
// The re-placement equals a fresh full layout of the same items
|
||||||
|
const { scroller: reference } = track(createScroller({ items }));
|
||||||
|
reference.refreshWithData(scroller.items.slice(), items.length, false);
|
||||||
|
expect(scroller.positions.map((p) => p.col)).toEqual(reference.positions.map((p) => p.col));
|
||||||
|
for (let i = 0; i < items.length; i++) {
|
||||||
|
expect(scroller.positions[i].top).toBeCloseTo(reference.positions[i].top);
|
||||||
|
}
|
||||||
|
|
||||||
|
// The rendered card was recreated in place with the update indicator
|
||||||
|
const updatedCard = grid.querySelector('.virtual-scroll-item.updated');
|
||||||
|
expect(updatedCard).not.toBeNull();
|
||||||
|
const indicator = updatedCard.querySelector('.update-indicator');
|
||||||
|
expect(indicator).not.toBeNull();
|
||||||
|
expect(indicator.textContent).toBe('Updated');
|
||||||
|
expect(updatedCard.querySelector('.card-preview').contains(indicator)).toBe(true);
|
||||||
|
expect(updatedCard.style.height).toBe(`${scroller.positions[2].height}px`);
|
||||||
|
expect(updatedCard.style.top).toBe(`${scroller.positions[2].top}px`);
|
||||||
|
});
|
||||||
|
|
||||||
|
it('updateSingleItem returns false for an unknown file path without throwing', () => {
|
||||||
|
const items = makeItems([{ width: 100, height: 100 }]);
|
||||||
|
const { scroller } = track(createScroller({ items }));
|
||||||
|
|
||||||
|
scroller.refreshWithData(items, items.length, false);
|
||||||
|
|
||||||
|
const warnSpy = vi.spyOn(console, 'warn').mockImplementation(() => {});
|
||||||
|
let result;
|
||||||
|
expect(() => {
|
||||||
|
result = scroller.updateSingleItem('/recipes/does-not-exist.png', { title: 'x' });
|
||||||
|
}).not.toThrow();
|
||||||
|
expect(result).toBe(false);
|
||||||
|
expect(warnSpy).toHaveBeenCalled();
|
||||||
|
warnSpy.mockRestore();
|
||||||
|
});
|
||||||
|
|
||||||
|
it('removeItemByFilePath re-places the remaining items and decrements the total', () => {
|
||||||
|
// Heights at ITEM_WIDTH=248: 496, 248, 124, 248, 248, 248
|
||||||
|
const items = makeItems([
|
||||||
|
{ width: 100, height: 200 },
|
||||||
|
{ width: 100, height: 100 },
|
||||||
|
{ width: 100, height: 50 },
|
||||||
|
{ width: 100, height: 100 },
|
||||||
|
{ width: 100, height: 100 },
|
||||||
|
{ width: 100, height: 100 },
|
||||||
|
]);
|
||||||
|
const { scroller } = track(createScroller({ items }));
|
||||||
|
|
||||||
|
scroller.refreshWithData(items, 60, false);
|
||||||
|
|
||||||
|
const result = scroller.removeItemByFilePath('/recipes/item-2.png');
|
||||||
|
|
||||||
|
expect(result).toBe(true);
|
||||||
|
expect(scroller.items.length).toBe(5);
|
||||||
|
expect(scroller.totalItems).toBe(59);
|
||||||
|
expect(scroller.positions.length).toBe(5);
|
||||||
|
|
||||||
|
// Remaining heights: 496, 248, 248, 248, 248 -> shortest-column placement
|
||||||
|
expect(scroller.positions.map((p) => p.col)).toEqual([0, 1, 2, 1, 2]);
|
||||||
|
expect(scroller.positions[3].top).toBeCloseTo(PAD_TOP + 248 + ROW_GAP); // 272
|
||||||
|
expect(scroller.positions[4].top).toBeCloseTo(PAD_TOP + 248 + ROW_GAP); // 272
|
||||||
|
|
||||||
|
// Spacer reflects the tallest remaining column
|
||||||
|
const maxColumnHeight = Math.max(...scroller.columnHeights);
|
||||||
|
const expected = maxColumnHeight - ROW_GAP + PAD_TOP + PAD_BOTTOM;
|
||||||
|
expect(scroller.spacerElement.style.height).toBe(`${expected}px`);
|
||||||
|
});
|
||||||
|
|
||||||
|
it('removeItemByFilePath returns false for an unknown file path', () => {
|
||||||
|
const items = makeItems([{ width: 100, height: 100 }]);
|
||||||
|
const { scroller } = track(createScroller({ items }));
|
||||||
|
|
||||||
|
scroller.refreshWithData(items, items.length, false);
|
||||||
|
|
||||||
|
const warnSpy = vi.spyOn(console, 'warn').mockImplementation(() => {});
|
||||||
|
expect(scroller.removeItemByFilePath('/recipes/missing.png')).toBe(false);
|
||||||
|
warnSpy.mockRestore();
|
||||||
|
});
|
||||||
|
|
||||||
|
it('removeMultipleItemsByFilePath re-places items with no layout gaps', () => {
|
||||||
|
const items = makeItems([
|
||||||
|
{ width: 100, height: 200 },
|
||||||
|
{ width: 100, height: 100 },
|
||||||
|
{ width: 100, height: 50 },
|
||||||
|
{ width: 100, height: 100 },
|
||||||
|
{ width: 100, height: 100 },
|
||||||
|
{ width: 100, height: 100 },
|
||||||
|
]);
|
||||||
|
const { scroller } = track(createScroller({ items }));
|
||||||
|
|
||||||
|
scroller.refreshWithData(items, 60, false);
|
||||||
|
|
||||||
|
const result = scroller.removeMultipleItemsByFilePath([
|
||||||
|
'/recipes/item-1.png',
|
||||||
|
'/recipes/item-3.png',
|
||||||
|
]);
|
||||||
|
|
||||||
|
expect(result).toBe(true);
|
||||||
|
expect(scroller.items.map((i) => i.file_path)).toEqual([
|
||||||
|
'/recipes/item-0.png',
|
||||||
|
'/recipes/item-2.png',
|
||||||
|
'/recipes/item-4.png',
|
||||||
|
'/recipes/item-5.png',
|
||||||
|
]);
|
||||||
|
expect(scroller.totalItems).toBe(58);
|
||||||
|
|
||||||
|
// The remaining items are laid out exactly as a fresh full placement:
|
||||||
|
// compare against a second scroller fed the same remaining items
|
||||||
|
const remaining = scroller.items.slice();
|
||||||
|
const { scroller: reference } = track(createScroller({ items: remaining }));
|
||||||
|
reference.refreshWithData(remaining, remaining.length, false);
|
||||||
|
|
||||||
|
expect(scroller.positions.map((p) => p.col)).toEqual(reference.positions.map((p) => p.col));
|
||||||
|
for (let i = 0; i < remaining.length; i++) {
|
||||||
|
expect(scroller.positions[i].top).toBeCloseTo(reference.positions[i].top);
|
||||||
|
expect(scroller.positions[i].left).toBeCloseTo(reference.positions[i].left);
|
||||||
|
}
|
||||||
|
});
|
||||||
|
|
||||||
|
it('removeMultipleItemsByFilePath returns false when nothing matches', () => {
|
||||||
|
const items = makeItems([{ width: 100, height: 100 }]);
|
||||||
|
const { scroller } = track(createScroller({ items }));
|
||||||
|
|
||||||
|
scroller.refreshWithData(items, items.length, false);
|
||||||
|
|
||||||
|
expect(scroller.removeMultipleItemsByFilePath(['/recipes/missing.png'])).toBe(false);
|
||||||
|
expect(scroller.removeMultipleItemsByFilePath([])).toBe(false);
|
||||||
|
});
|
||||||
|
|
||||||
|
it('disable stops rendering and enable recreates the spacer after innerHTML is cleared', async () => {
|
||||||
|
const items = makeItems([
|
||||||
|
{ width: 100, height: 200 },
|
||||||
|
{ width: 100, height: 100 },
|
||||||
|
{ width: 100, height: 50 },
|
||||||
|
]);
|
||||||
|
const { scroller, grid, wrapper } = track(createScroller({ items, viewportHeight: 3000 }));
|
||||||
|
|
||||||
|
scroller.refreshWithData(items, items.length, false);
|
||||||
|
wrapper.scrollTop = 0;
|
||||||
|
scroller.overscan = 5;
|
||||||
|
scroller.renderItems();
|
||||||
|
expect(grid.querySelectorAll('.virtual-scroll-item').length).toBe(3);
|
||||||
|
|
||||||
|
scroller.disable();
|
||||||
|
|
||||||
|
expect(scroller.disabled).toBe(true);
|
||||||
|
expect(grid.querySelectorAll('.virtual-scroll-item').length).toBe(0);
|
||||||
|
expect(scroller.spacerElement.style.display).toBe('none');
|
||||||
|
|
||||||
|
// Duplicates mode wipes the grid contents, destroying the spacer
|
||||||
|
grid.innerHTML = '';
|
||||||
|
expect(grid.contains(scroller.spacerElement)).toBe(false);
|
||||||
|
|
||||||
|
scroller.enable();
|
||||||
|
|
||||||
|
expect(scroller.disabled).toBe(false);
|
||||||
|
expect(grid.contains(scroller.spacerElement)).toBe(true);
|
||||||
|
expect(scroller.spacerElement.className).toBe('virtual-scroll-spacer');
|
||||||
|
|
||||||
|
// Full re-placement ran synchronously on re-enable
|
||||||
|
expect(scroller.positions.length).toBe(items.length);
|
||||||
|
|
||||||
|
// Rendering resumes after the scheduled rAF
|
||||||
|
await new Promise((resolve) => setTimeout(resolve, 20));
|
||||||
|
expect(grid.querySelectorAll('.virtual-scroll-item').length).toBe(3);
|
||||||
|
});
|
||||||
|
|
||||||
|
it('getAdjacentItemByFilePath loads more pages when the target is beyond loaded items', async () => {
|
||||||
|
const page1 = [0, 1, 2].map((i) => ({
|
||||||
|
file_path: `/recipes/page1-${i}.png`,
|
||||||
|
width: 100,
|
||||||
|
height: 100,
|
||||||
|
}));
|
||||||
|
const page2 = [0, 1].map((i) => ({
|
||||||
|
file_path: `/recipes/page2-${i}.png`,
|
||||||
|
width: 100,
|
||||||
|
height: 100,
|
||||||
|
}));
|
||||||
|
const fetchMock = vi.fn(async () => ({ items: page2, totalItems: 5, hasMore: false }));
|
||||||
|
const { scroller } = track(createScroller({ fetchItemsFn: fetchMock }));
|
||||||
|
|
||||||
|
scroller.refreshWithData(page1, 5, true);
|
||||||
|
const pageState = getCurrentPageState();
|
||||||
|
const expectedPage = pageState.currentPage;
|
||||||
|
|
||||||
|
const result = await scroller.getAdjacentItemByFilePath('/recipes/page1-2.png', 'next');
|
||||||
|
|
||||||
|
expect(fetchMock).toHaveBeenCalledWith(expectedPage, scroller.pageSize);
|
||||||
|
expect(result).not.toBeNull();
|
||||||
|
expect(result.index).toBe(3);
|
||||||
|
expect(result.item.file_path).toBe('/recipes/page2-0.png');
|
||||||
|
});
|
||||||
|
|
||||||
|
it('getAdjacentItemByFilePath returns null at boundaries and for unknown paths', async () => {
|
||||||
|
const items = makeItems([{ width: 100, height: 100 }, { width: 100, height: 100 }]);
|
||||||
|
const { scroller } = track(createScroller({ items }));
|
||||||
|
|
||||||
|
scroller.refreshWithData(items, items.length, false);
|
||||||
|
|
||||||
|
await expect(scroller.getAdjacentItemByFilePath('/recipes/item-0.png', 'prev')).resolves.toBeNull();
|
||||||
|
await expect(scroller.getAdjacentItemByFilePath('/recipes/item-1.png', 'next')).resolves.toBeNull();
|
||||||
|
await expect(scroller.getAdjacentItemByFilePath('/recipes/missing.png', 'next')).resolves.toBeNull();
|
||||||
|
});
|
||||||
|
|
||||||
|
it('getNavigationState reports index, prev/next availability and totals', () => {
|
||||||
|
const items = makeItems([{ width: 100, height: 100 }, { width: 100, height: 100 }]);
|
||||||
|
const { scroller } = track(createScroller({ items }));
|
||||||
|
|
||||||
|
scroller.refreshWithData(items, 10, true);
|
||||||
|
|
||||||
|
expect(scroller.getNavigationState('/recipes/item-0.png')).toEqual({
|
||||||
|
index: 0,
|
||||||
|
hasPrev: false,
|
||||||
|
hasNext: true,
|
||||||
|
loadedItems: 2,
|
||||||
|
totalItems: 10,
|
||||||
|
});
|
||||||
|
|
||||||
|
const last = scroller.getNavigationState('/recipes/item-1.png');
|
||||||
|
expect(last.index).toBe(1);
|
||||||
|
expect(last.hasPrev).toBe(true);
|
||||||
|
// hasMore keeps forward navigation available past the loaded window
|
||||||
|
expect(last.hasNext).toBe(true);
|
||||||
|
|
||||||
|
expect(scroller.getNavigationState('/recipes/missing.png').index).toBe(-1);
|
||||||
|
expect(scroller.findIndexByFilePath('/recipes/item-1.png')).toBe(1);
|
||||||
|
expect(scroller.findIndexByFilePath('')).toBe(-1);
|
||||||
|
});
|
||||||
|
});
|
||||||
@@ -1,4 +1,11 @@
|
|||||||
from py.nodes.lora_stack_combiner import LoraStackCombinerLM
|
import types
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
|
||||||
|
from py.nodes.lora_stack_combiner import (
|
||||||
|
LoraStackCombinerLM,
|
||||||
|
_LoraStackOptionalInputs,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
def test_combine_stacks_preserves_order():
|
def test_combine_stacks_preserves_order():
|
||||||
@@ -49,3 +56,109 @@ def test_combine_stacks_allows_duplicate_entries():
|
|||||||
(combined_stack,) = node.combine_stacks([duplicate_entry], [duplicate_entry])
|
(combined_stack,) = node.combine_stacks([duplicate_entry], [duplicate_entry])
|
||||||
|
|
||||||
assert combined_stack == [duplicate_entry, duplicate_entry]
|
assert combined_stack == [duplicate_entry, duplicate_entry]
|
||||||
|
|
||||||
|
|
||||||
|
def test_combine_stacks_returns_empty_when_both_unconnected():
|
||||||
|
node = LoraStackCombinerLM()
|
||||||
|
|
||||||
|
(combined_stack,) = node.combine_stacks()
|
||||||
|
|
||||||
|
assert combined_stack == []
|
||||||
|
|
||||||
|
|
||||||
|
def test_combine_stacks_returns_other_when_one_unconnected():
|
||||||
|
node = LoraStackCombinerLM()
|
||||||
|
stack_a = [("folder/a.safetensors", 0.7, 0.6)]
|
||||||
|
|
||||||
|
(combined_stack_a,) = node.combine_stacks(lora_stack1=stack_a)
|
||||||
|
(combined_stack_b,) = node.combine_stacks(lora_stack2=stack_a)
|
||||||
|
|
||||||
|
assert combined_stack_a == stack_a
|
||||||
|
assert combined_stack_b == stack_a
|
||||||
|
|
||||||
|
|
||||||
|
def test_combine_stacks_with_dynamic_third_slot():
|
||||||
|
node = LoraStackCombinerLM()
|
||||||
|
stack_a = [("folder/a.safetensors", 0.7, 0.6)]
|
||||||
|
stack_b = [("folder/b.safetensors", 0.8, 0.8)]
|
||||||
|
stack_c = [("folder/c.safetensors", 1.0, 0.9)]
|
||||||
|
|
||||||
|
(combined_stack,) = node.combine_stacks(
|
||||||
|
lora_stack1=stack_a, lora_stack2=stack_b, lora_stack3=stack_c
|
||||||
|
)
|
||||||
|
|
||||||
|
assert combined_stack == stack_a + stack_b + stack_c
|
||||||
|
|
||||||
|
|
||||||
|
def test_combine_stacks_orders_by_slot_number_not_call_order():
|
||||||
|
node = LoraStackCombinerLM()
|
||||||
|
stack_a = [("folder/a.safetensors", 0.7, 0.6)]
|
||||||
|
stack_b = [("folder/b.safetensors", 0.8, 0.8)]
|
||||||
|
stack_c = [("folder/c.safetensors", 1.0, 0.9)]
|
||||||
|
|
||||||
|
(combined_stack,) = node.combine_stacks(
|
||||||
|
lora_stack3=stack_c, lora_stack2=stack_b, lora_stack1=stack_a
|
||||||
|
)
|
||||||
|
|
||||||
|
assert combined_stack == stack_a + stack_b + stack_c
|
||||||
|
|
||||||
|
|
||||||
|
def test_combine_stacks_accepts_only_dynamic_slot():
|
||||||
|
node = LoraStackCombinerLM()
|
||||||
|
stack_c = [("folder/c.safetensors", 1.0, 0.9)]
|
||||||
|
|
||||||
|
(combined_stack,) = node.combine_stacks(lora_stack3=stack_c)
|
||||||
|
|
||||||
|
assert combined_stack == stack_c
|
||||||
|
|
||||||
|
|
||||||
|
def test_combine_stacks_handles_legacy_input_names():
|
||||||
|
node = LoraStackCombinerLM()
|
||||||
|
stack_a = [("folder/a.safetensors", 0.7, 0.6)]
|
||||||
|
stack_b = [("folder/b.safetensors", 0.8, 0.8)]
|
||||||
|
|
||||||
|
(combined_stack,) = node.combine_stacks(lora_stack_a=stack_a, lora_stack_b=stack_b)
|
||||||
|
|
||||||
|
assert combined_stack == stack_a + stack_b
|
||||||
|
|
||||||
|
|
||||||
|
def test_input_types_exposes_two_default_slots():
|
||||||
|
input_types = LoraStackCombinerLM.INPUT_TYPES()
|
||||||
|
|
||||||
|
assert set(input_types["optional"]) == {"lora_stack1", "lora_stack2"}
|
||||||
|
assert input_types["optional"]["lora_stack1"][0] == "LORA_STACK"
|
||||||
|
assert input_types["optional"]["lora_stack2"][0] == "LORA_STACK"
|
||||||
|
|
||||||
|
|
||||||
|
def test_input_types_recognizes_dynamic_slots_from_get_input_info(monkeypatch):
|
||||||
|
frames = [None, None, types.SimpleNamespace(function="get_input_info")]
|
||||||
|
monkeypatch.setattr(
|
||||||
|
"py.nodes.lora_stack_combiner.inspect.stack", lambda: frames
|
||||||
|
)
|
||||||
|
|
||||||
|
input_types = LoraStackCombinerLM.INPUT_TYPES()
|
||||||
|
optional = input_types["optional"]
|
||||||
|
|
||||||
|
assert "lora_stack3" in optional
|
||||||
|
assert optional["lora_stack3"][0] == "LORA_STACK"
|
||||||
|
assert "lora_stack25" in optional
|
||||||
|
assert optional["lora_stack25"][0] == "LORA_STACK"
|
||||||
|
|
||||||
|
|
||||||
|
def test_lora_stack_optional_inputs_proxy():
|
||||||
|
proxy = _LoraStackOptionalInputs({"lora_stack1": ("LORA_STACK", {})})
|
||||||
|
|
||||||
|
assert "lora_stack1" in proxy
|
||||||
|
assert "lora_stack2" in proxy
|
||||||
|
assert "lora_stack10" in proxy
|
||||||
|
assert "lora_stack_a" in proxy
|
||||||
|
assert "lora_stack" not in proxy
|
||||||
|
assert "lora_stacka" not in proxy
|
||||||
|
assert "lora_stack_1" not in proxy
|
||||||
|
assert "text" not in proxy
|
||||||
|
|
||||||
|
assert proxy["lora_stack1"][0] == "LORA_STACK"
|
||||||
|
assert proxy["lora_stack5"][0] == "LORA_STACK"
|
||||||
|
|
||||||
|
with pytest.raises(KeyError):
|
||||||
|
proxy["not_a_stack"]
|
||||||
|
|||||||
@@ -86,6 +86,41 @@ def test_save_image_skips_png_parameters_when_metadata_disabled_and_keeps_workfl
|
|||||||
assert img.info["workflow"] == json.dumps(workflow)
|
assert img.info["workflow"] == json.dumps(workflow)
|
||||||
|
|
||||||
|
|
||||||
|
def test_save_image_does_not_append_loras_to_prompt_by_default(monkeypatch, tmp_path):
|
||||||
|
_configure_save_paths(monkeypatch, tmp_path)
|
||||||
|
_configure_metadata(
|
||||||
|
monkeypatch,
|
||||||
|
{"prompt": "prompt text", "seed": 123, "loras": "<lora:foo:0.7>"},
|
||||||
|
)
|
||||||
|
|
||||||
|
node = SaveImageLM()
|
||||||
|
node.save_images([_make_image()], "ComfyUI", "png", id="node-1")
|
||||||
|
|
||||||
|
image_path = tmp_path / "sample_00001_.png"
|
||||||
|
with Image.open(image_path) as img:
|
||||||
|
assert "<lora:" not in img.info["parameters"]
|
||||||
|
assert img.info["parameters"] == "prompt text\nSeed: 123, Version: ComfyUI"
|
||||||
|
|
||||||
|
|
||||||
|
def test_save_image_appends_loras_to_prompt_when_enabled(monkeypatch, tmp_path):
|
||||||
|
_configure_save_paths(monkeypatch, tmp_path)
|
||||||
|
_configure_metadata(
|
||||||
|
monkeypatch,
|
||||||
|
{"prompt": "prompt text", "seed": 123, "loras": "<lora:foo:0.7>"},
|
||||||
|
)
|
||||||
|
|
||||||
|
node = SaveImageLM()
|
||||||
|
node.save_images(
|
||||||
|
[_make_image()], "ComfyUI", "png", id="node-1", add_loras_to_prompt=True
|
||||||
|
)
|
||||||
|
|
||||||
|
image_path = tmp_path / "sample_00001_.png"
|
||||||
|
with Image.open(image_path) as img:
|
||||||
|
assert img.info["parameters"] == (
|
||||||
|
"prompt text\n<lora:foo:0.7>\nSeed: 123, Version: ComfyUI"
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
def test_save_image_skips_jpeg_metadata_when_disabled(monkeypatch, tmp_path):
|
def test_save_image_skips_jpeg_metadata_when_disabled(monkeypatch, tmp_path):
|
||||||
_configure_save_paths(monkeypatch, tmp_path)
|
_configure_save_paths(monkeypatch, tmp_path)
|
||||||
_configure_metadata(monkeypatch, {"prompt": "prompt text", "seed": 123})
|
_configure_metadata(monkeypatch, {"prompt": "prompt text", "seed": 123})
|
||||||
@@ -451,6 +486,14 @@ class TestParameterDefaultConsistency:
|
|||||||
assert SaveImageLM.save_images.__defaults__[5] == 0
|
assert SaveImageLM.save_images.__defaults__[5] == 0
|
||||||
assert SaveImageLM.process_image.__defaults__[7] == 0
|
assert SaveImageLM.process_image.__defaults__[7] == 0
|
||||||
|
|
||||||
|
def test_add_loras_to_prompt_defaults_are_consistent(self):
|
||||||
|
input_types = SaveImageLM.INPUT_TYPES()
|
||||||
|
optional = input_types["optional"]
|
||||||
|
|
||||||
|
assert optional["add_loras_to_prompt"][1]["default"] is False
|
||||||
|
assert SaveImageLM.save_images.__defaults__[-1] is False
|
||||||
|
assert SaveImageLM.process_image.__defaults__[-1] is False
|
||||||
|
|
||||||
|
|
||||||
def test_png_does_not_pass_webp_method_or_jpeg_subsampling(monkeypatch, tmp_path):
|
def test_png_does_not_pass_webp_method_or_jpeg_subsampling(monkeypatch, tmp_path):
|
||||||
_configure_save_paths(monkeypatch, tmp_path)
|
_configure_save_paths(monkeypatch, tmp_path)
|
||||||
|
|||||||
@@ -900,18 +900,28 @@ class FakeMetadataProvider:
|
|||||||
async def get_model_versions(self, _model_id):
|
async def get_model_versions(self, _model_id):
|
||||||
return {"modelVersions": [], "name": "", "type": "lora"}
|
return {"modelVersions": [], "name": "", "type": "lora"}
|
||||||
|
|
||||||
async def get_user_models(self, _username):
|
async def get_user_models(self, _username, cursor=None):
|
||||||
return []
|
return {"items": [], "nextCursor": None}
|
||||||
|
|
||||||
|
async def get_creator_model_count(self, _username):
|
||||||
|
return None
|
||||||
|
|
||||||
|
|
||||||
class FakeUserModelsProvider(FakeMetadataProvider):
|
class FakeUserModelsProvider(FakeMetadataProvider):
|
||||||
def __init__(self, models):
|
def __init__(self, models, next_cursor=None, estimated_total=None):
|
||||||
self.models = models
|
self.models = models
|
||||||
|
self.next_cursor = next_cursor
|
||||||
|
self.estimated_total = estimated_total
|
||||||
self.received_usernames: list[str] = []
|
self.received_usernames: list[str] = []
|
||||||
|
self.received_cursors: list = []
|
||||||
|
|
||||||
async def get_user_models(self, username):
|
async def get_user_models(self, username, cursor=None):
|
||||||
self.received_usernames.append(username)
|
self.received_usernames.append(username)
|
||||||
return self.models
|
self.received_cursors.append(cursor)
|
||||||
|
return {"items": self.models, "nextCursor": self.next_cursor}
|
||||||
|
|
||||||
|
async def get_creator_model_count(self, _username):
|
||||||
|
return self.estimated_total
|
||||||
|
|
||||||
|
|
||||||
async def fake_metadata_provider_factory():
|
async def fake_metadata_provider_factory():
|
||||||
@@ -1286,6 +1296,88 @@ async def test_get_civitai_user_models_requires_username():
|
|||||||
assert "username" in payload["error"].lower()
|
assert "username" in payload["error"].lower()
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_get_civitai_user_models_returns_pagination_fields():
|
||||||
|
models = [
|
||||||
|
{
|
||||||
|
"id": 1,
|
||||||
|
"name": "Model A",
|
||||||
|
"type": "LORA",
|
||||||
|
"tags": [],
|
||||||
|
"modelVersions": [
|
||||||
|
{"id": 100, "name": "v1", "images": [{"url": "http://example.com/a.jpg"}]},
|
||||||
|
],
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"id": 2,
|
||||||
|
"name": "Unsupported",
|
||||||
|
"type": "Other",
|
||||||
|
"modelVersions": [{"id": 200, "name": "v1"}],
|
||||||
|
},
|
||||||
|
]
|
||||||
|
|
||||||
|
provider = FakeUserModelsProvider(models, next_cursor="cursor-token", estimated_total=2140)
|
||||||
|
|
||||||
|
async def provider_factory():
|
||||||
|
return provider
|
||||||
|
|
||||||
|
handler = ModelLibraryHandler(
|
||||||
|
ServiceRegistryAdapter(
|
||||||
|
get_lora_scanner=fake_scanner_factory,
|
||||||
|
get_checkpoint_scanner=fake_scanner_factory,
|
||||||
|
get_embedding_scanner=fake_scanner_factory,
|
||||||
|
get_downloaded_version_history_service=fake_download_history_service_factory,
|
||||||
|
),
|
||||||
|
metadata_provider_factory=provider_factory,
|
||||||
|
)
|
||||||
|
|
||||||
|
response = await handler.get_civitai_user_models(
|
||||||
|
FakeRequest(query={"username": "pixel"})
|
||||||
|
)
|
||||||
|
payload = json.loads(response.text)
|
||||||
|
|
||||||
|
assert response.status == 200
|
||||||
|
assert payload["success"] is True
|
||||||
|
# modelCount only counts models surviving the type filter
|
||||||
|
assert payload["modelCount"] == 1
|
||||||
|
assert payload["nextCursor"] == "cursor-token"
|
||||||
|
assert payload["hasMore"] is True
|
||||||
|
# first page includes the estimated total
|
||||||
|
assert payload["estimatedTotal"] == 2140
|
||||||
|
assert provider.received_cursors == [None]
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_get_civitai_user_models_passes_cursor_and_omits_estimate():
|
||||||
|
provider = FakeUserModelsProvider([], next_cursor=None, estimated_total=999)
|
||||||
|
|
||||||
|
async def provider_factory():
|
||||||
|
return provider
|
||||||
|
|
||||||
|
handler = ModelLibraryHandler(
|
||||||
|
ServiceRegistryAdapter(
|
||||||
|
get_lora_scanner=fake_scanner_factory,
|
||||||
|
get_checkpoint_scanner=fake_scanner_factory,
|
||||||
|
get_embedding_scanner=fake_scanner_factory,
|
||||||
|
get_downloaded_version_history_service=fake_download_history_service_factory,
|
||||||
|
),
|
||||||
|
metadata_provider_factory=provider_factory,
|
||||||
|
)
|
||||||
|
|
||||||
|
response = await handler.get_civitai_user_models(
|
||||||
|
FakeRequest(query={"username": "pixel", "cursor": "opaque-token"})
|
||||||
|
)
|
||||||
|
payload = json.loads(response.text)
|
||||||
|
|
||||||
|
assert response.status == 200
|
||||||
|
assert payload["success"] is True
|
||||||
|
assert payload["nextCursor"] is None
|
||||||
|
assert payload["hasMore"] is False
|
||||||
|
# cursor requests must not include the estimated total
|
||||||
|
assert payload["estimatedTotal"] is None
|
||||||
|
assert provider.received_cursors == ["opaque-token"]
|
||||||
|
|
||||||
|
|
||||||
def test_ensure_handler_mapping_caches_result():
|
def test_ensure_handler_mapping_caches_result():
|
||||||
call_records = []
|
call_records = []
|
||||||
|
|
||||||
|
|||||||
@@ -2,6 +2,7 @@
|
|||||||
|
|
||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import asyncio
|
||||||
import json
|
import json
|
||||||
from contextlib import asynccontextmanager
|
from contextlib import asynccontextmanager
|
||||||
from dataclasses import dataclass
|
from dataclasses import dataclass
|
||||||
@@ -11,6 +12,7 @@ from typing import Any, AsyncIterator, Dict, List, Optional
|
|||||||
|
|
||||||
from aiohttp import FormData, web
|
from aiohttp import FormData, web
|
||||||
from aiohttp.test_utils import TestClient, TestServer
|
from aiohttp.test_utils import TestClient, TestServer
|
||||||
|
from PIL import Image
|
||||||
|
|
||||||
from py.config import config
|
from py.config import config
|
||||||
from py.routes import base_recipe_routes
|
from py.routes import base_recipe_routes
|
||||||
@@ -368,6 +370,163 @@ async def test_list_recipes_provides_file_urls(monkeypatch, tmp_path: Path) -> N
|
|||||||
assert payload["items"][0]["loras"] == []
|
assert payload["items"][0]["loras"] == []
|
||||||
|
|
||||||
|
|
||||||
|
async def test_list_recipes_exposes_preview_dimensions(
|
||||||
|
monkeypatch, tmp_path: Path
|
||||||
|
) -> None:
|
||||||
|
"""(a) Image recipe items carry integer width/height from the on-disk file."""
|
||||||
|
async with recipe_harness(monkeypatch, tmp_path) as harness:
|
||||||
|
recipe_path = harness.tmp_dir / "recipes" / "real.png"
|
||||||
|
recipe_path.parent.mkdir(parents=True, exist_ok=True)
|
||||||
|
Image.new("RGB", (64, 32), color="red").save(recipe_path)
|
||||||
|
|
||||||
|
harness.scanner.listing_items = [
|
||||||
|
{
|
||||||
|
"id": "recipe-1",
|
||||||
|
"file_path": str(recipe_path),
|
||||||
|
"title": "Image Recipe",
|
||||||
|
"loras": [],
|
||||||
|
}
|
||||||
|
]
|
||||||
|
harness.scanner.cached_raw = list(harness.scanner.listing_items)
|
||||||
|
|
||||||
|
response = await harness.client.get("/api/lm/recipes")
|
||||||
|
payload = await response.json()
|
||||||
|
|
||||||
|
assert response.status == 200
|
||||||
|
item = payload["items"][0]
|
||||||
|
assert item["width"] == 64
|
||||||
|
assert item["height"] == 32
|
||||||
|
assert isinstance(item["width"], int)
|
||||||
|
assert isinstance(item["height"], int)
|
||||||
|
|
||||||
|
|
||||||
|
async def test_list_recipes_omits_dimensions_for_video_and_missing(
|
||||||
|
monkeypatch, tmp_path: Path
|
||||||
|
) -> None:
|
||||||
|
"""(b) Video/missing-image recipes omit width/height yet still return 200."""
|
||||||
|
async with recipe_harness(monkeypatch, tmp_path) as harness:
|
||||||
|
harness.scanner.listing_items = [
|
||||||
|
{
|
||||||
|
"id": "recipe-video",
|
||||||
|
"file_path": str(harness.tmp_dir / "recipes" / "preview.mp4"),
|
||||||
|
"title": "Video Recipe",
|
||||||
|
"loras": [],
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"id": "recipe-missing",
|
||||||
|
"file_path": str(harness.tmp_dir / "recipes" / "gone.png"),
|
||||||
|
"title": "Missing Recipe",
|
||||||
|
"loras": [],
|
||||||
|
},
|
||||||
|
]
|
||||||
|
harness.scanner.cached_raw = list(harness.scanner.listing_items)
|
||||||
|
|
||||||
|
response = await harness.client.get("/api/lm/recipes")
|
||||||
|
payload = await response.json()
|
||||||
|
|
||||||
|
assert response.status == 200
|
||||||
|
for item in payload["items"]:
|
||||||
|
assert "width" not in item
|
||||||
|
assert "height" not in item
|
||||||
|
|
||||||
|
|
||||||
|
async def test_list_recipes_offloads_dimensions_to_thread(
|
||||||
|
monkeypatch, tmp_path: Path
|
||||||
|
) -> None:
|
||||||
|
"""(c) Dimension reads run through asyncio.to_thread for every item."""
|
||||||
|
async with recipe_harness(monkeypatch, tmp_path) as harness:
|
||||||
|
recipe_path = harness.tmp_dir / "recipes" / "real.png"
|
||||||
|
recipe_path.parent.mkdir(parents=True, exist_ok=True)
|
||||||
|
Image.new("RGB", (16, 48), color="blue").save(recipe_path)
|
||||||
|
|
||||||
|
harness.scanner.listing_items = [
|
||||||
|
{
|
||||||
|
"id": "recipe-1",
|
||||||
|
"file_path": str(recipe_path),
|
||||||
|
"title": "Image Recipe",
|
||||||
|
"loras": [],
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"id": "recipe-2",
|
||||||
|
"file_path": str(harness.tmp_dir / "recipes" / "gone.png"),
|
||||||
|
"title": "Missing Recipe",
|
||||||
|
"loras": [],
|
||||||
|
},
|
||||||
|
]
|
||||||
|
harness.scanner.cached_raw = list(harness.scanner.listing_items)
|
||||||
|
|
||||||
|
real_to_thread = asyncio.to_thread
|
||||||
|
to_thread_calls: list[tuple] = []
|
||||||
|
|
||||||
|
async def counting_to_thread(fn, *args, **kwargs):
|
||||||
|
to_thread_calls.append((fn, args, kwargs))
|
||||||
|
return await real_to_thread(fn, *args, **kwargs)
|
||||||
|
|
||||||
|
monkeypatch.setattr(asyncio, "to_thread", counting_to_thread)
|
||||||
|
|
||||||
|
response = await harness.client.get("/api/lm/recipes")
|
||||||
|
payload = await response.json()
|
||||||
|
|
||||||
|
assert response.status == 200
|
||||||
|
assert len(to_thread_calls) >= len(harness.scanner.listing_items)
|
||||||
|
assert payload["items"][0]["width"] == 16
|
||||||
|
assert payload["items"][0]["height"] == 48
|
||||||
|
assert "width" not in payload["items"][1]
|
||||||
|
assert "height" not in payload["items"][1]
|
||||||
|
|
||||||
|
|
||||||
|
async def test_list_recipes_batches_dimensions_for_mixed_items(
|
||||||
|
monkeypatch, tmp_path: Path
|
||||||
|
) -> None:
|
||||||
|
"""(d) Mixed file_path presence: dims align per-item with their own files."""
|
||||||
|
async with recipe_harness(monkeypatch, tmp_path) as harness:
|
||||||
|
wide_path = harness.tmp_dir / "recipes" / "wide.png"
|
||||||
|
tall_path = harness.tmp_dir / "recipes" / "tall.png"
|
||||||
|
wide_path.parent.mkdir(parents=True, exist_ok=True)
|
||||||
|
Image.new("RGB", (120, 40), color="green").save(wide_path)
|
||||||
|
Image.new("RGB", (30, 90), color="blue").save(tall_path)
|
||||||
|
|
||||||
|
harness.scanner.listing_items = [
|
||||||
|
{
|
||||||
|
"id": "recipe-wide",
|
||||||
|
"file_path": str(wide_path),
|
||||||
|
"title": "Wide",
|
||||||
|
"loras": [],
|
||||||
|
},
|
||||||
|
{"id": "recipe-none", "title": "No Preview", "loras": []},
|
||||||
|
{
|
||||||
|
"id": "recipe-tall",
|
||||||
|
"file_path": str(tall_path),
|
||||||
|
"title": "Tall",
|
||||||
|
"loras": [],
|
||||||
|
},
|
||||||
|
{"id": "recipe-none-2", "title": "No Preview 2", "loras": []},
|
||||||
|
]
|
||||||
|
harness.scanner.cached_raw = list(harness.scanner.listing_items)
|
||||||
|
|
||||||
|
response = await harness.client.get("/api/lm/recipes")
|
||||||
|
payload = await response.json()
|
||||||
|
|
||||||
|
assert response.status == 200
|
||||||
|
items = payload["items"]
|
||||||
|
|
||||||
|
# Items with a file_path carry integer dims read from their own file;
|
||||||
|
# the two images differ in both dimensions so a shifted pair would
|
||||||
|
# fail these assertions.
|
||||||
|
assert items[0]["width"] == 120
|
||||||
|
assert items[0]["height"] == 40
|
||||||
|
assert isinstance(items[0]["width"], int)
|
||||||
|
assert isinstance(items[0]["height"], int)
|
||||||
|
assert items[2]["width"] == 30
|
||||||
|
assert items[2]["height"] == 90
|
||||||
|
|
||||||
|
# Items without a file_path get the no-preview fallback and omit dims.
|
||||||
|
for item in (items[1], items[3]):
|
||||||
|
assert "width" not in item
|
||||||
|
assert "height" not in item
|
||||||
|
assert item["file_url"] == "/loras_static/images/no-preview.png"
|
||||||
|
|
||||||
|
|
||||||
async def test_list_recipes_passes_checkpoint_hash_filter(
|
async def test_list_recipes_passes_checkpoint_hash_filter(
|
||||||
monkeypatch, tmp_path: Path
|
monkeypatch, tmp_path: Path
|
||||||
) -> None:
|
) -> None:
|
||||||
@@ -909,8 +1068,6 @@ async def test_batch_import_start_missing_source(monkeypatch, tmp_path: Path) ->
|
|||||||
|
|
||||||
|
|
||||||
async def test_batch_import_start_already_running(monkeypatch, tmp_path: Path) -> None:
|
async def test_batch_import_start_already_running(monkeypatch, tmp_path: Path) -> None:
|
||||||
import asyncio
|
|
||||||
|
|
||||||
async with recipe_harness(monkeypatch, tmp_path) as harness:
|
async with recipe_harness(monkeypatch, tmp_path) as harness:
|
||||||
original_analyze = harness.analysis.analyze_remote_image
|
original_analyze = harness.analysis.analyze_remote_image
|
||||||
|
|
||||||
|
|||||||
@@ -183,7 +183,7 @@ class FakeCache:
|
|||||||
def __init__(self, items):
|
def __init__(self, items):
|
||||||
self.items = list(items)
|
self.items = list(items)
|
||||||
|
|
||||||
async def get_sorted_data(self, sort_key, order):
|
async def get_sorted_data(self, sort_key, order, seed=None):
|
||||||
if sort_key == "name":
|
if sort_key == "name":
|
||||||
data = sorted(self.items, key=lambda x: x["model_name"].lower())
|
data = sorted(self.items, key=lambda x: x["model_name"].lower())
|
||||||
if order == "desc":
|
if order == "desc":
|
||||||
|
|||||||
@@ -363,6 +363,148 @@ async def test_check_pending_models_handles_corrupted_progress_file(
|
|||||||
assert result["pending_count"] == 1
|
assert result["pending_count"] == 1
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
@pytest.mark.usefixtures("tmp_path")
|
||||||
|
async def test_check_pending_models_uses_bulk_folder_index_for_large_libraries(
|
||||||
|
monkeypatch: pytest.MonkeyPatch,
|
||||||
|
tmp_path,
|
||||||
|
settings_manager,
|
||||||
|
):
|
||||||
|
"""For >1000 candidates the pre-check scans the library root once instead of
|
||||||
|
probing every folder individually."""
|
||||||
|
|
||||||
|
ws_manager = RecordingWebSocketManager()
|
||||||
|
manager = download_module.DownloadManager(ws_manager=ws_manager)
|
||||||
|
|
||||||
|
monkeypatch.setitem(settings_manager.settings, "example_images_path", str(tmp_path))
|
||||||
|
|
||||||
|
# 1500 unprocessed models triggers the bulk lookup path
|
||||||
|
models = [
|
||||||
|
{"sha256": f"{i:064x}", "model_name": f"Model {i}"}
|
||||||
|
for i in range(1500)
|
||||||
|
]
|
||||||
|
|
||||||
|
# Create folders with files for the first 500 models
|
||||||
|
for i in range(500):
|
||||||
|
model_dir = tmp_path / f"{i:064x}"
|
||||||
|
model_dir.mkdir()
|
||||||
|
(model_dir / "image_0.png").write_text("data")
|
||||||
|
|
||||||
|
_patch_scanners(monkeypatch, lora_scanner=StubScanner(models))
|
||||||
|
|
||||||
|
per_model_checks = 0
|
||||||
|
|
||||||
|
def counting_model_directory_has_files(path: str) -> bool:
|
||||||
|
nonlocal per_model_checks
|
||||||
|
per_model_checks += 1
|
||||||
|
return False
|
||||||
|
|
||||||
|
monkeypatch.setattr(
|
||||||
|
download_module,
|
||||||
|
"_model_directory_has_files",
|
||||||
|
counting_model_directory_has_files,
|
||||||
|
)
|
||||||
|
|
||||||
|
result = await manager.check_pending_models(["lora"])
|
||||||
|
|
||||||
|
assert result["success"] is True
|
||||||
|
assert result["total_models"] == 1500
|
||||||
|
assert result["pending_count"] == 1000
|
||||||
|
assert result["needs_download"] is True
|
||||||
|
# The per-folder check should not be used once we cross the threshold.
|
||||||
|
assert per_model_checks == 0
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
@pytest.mark.usefixtures("tmp_path")
|
||||||
|
async def test_check_pending_models_uses_per_folder_check_for_small_candidate_sets(
|
||||||
|
monkeypatch: pytest.MonkeyPatch,
|
||||||
|
tmp_path,
|
||||||
|
settings_manager,
|
||||||
|
):
|
||||||
|
"""For <=1000 candidates the pre-check keeps the accurate per-folder path."""
|
||||||
|
|
||||||
|
ws_manager = RecordingWebSocketManager()
|
||||||
|
manager = download_module.DownloadManager(ws_manager=ws_manager)
|
||||||
|
|
||||||
|
monkeypatch.setitem(settings_manager.settings, "example_images_path", str(tmp_path))
|
||||||
|
|
||||||
|
models = [
|
||||||
|
{"sha256": f"{i:064x}", "model_name": f"Model {i}"}
|
||||||
|
for i in range(500)
|
||||||
|
]
|
||||||
|
|
||||||
|
# Create folders with files for the first 200 models
|
||||||
|
for i in range(200):
|
||||||
|
model_dir = tmp_path / f"{i:064x}"
|
||||||
|
model_dir.mkdir()
|
||||||
|
(model_dir / "image_0.png").write_text("data")
|
||||||
|
|
||||||
|
_patch_scanners(monkeypatch, lora_scanner=StubScanner(models))
|
||||||
|
|
||||||
|
per_model_checks = 0
|
||||||
|
original_has_files = download_module._model_directory_has_files
|
||||||
|
|
||||||
|
def counting_model_directory_has_files(path: str) -> bool:
|
||||||
|
nonlocal per_model_checks
|
||||||
|
per_model_checks += 1
|
||||||
|
return original_has_files(path)
|
||||||
|
|
||||||
|
monkeypatch.setattr(
|
||||||
|
download_module,
|
||||||
|
"_model_directory_has_files",
|
||||||
|
counting_model_directory_has_files,
|
||||||
|
)
|
||||||
|
|
||||||
|
result = await manager.check_pending_models(["lora"])
|
||||||
|
|
||||||
|
assert result["success"] is True
|
||||||
|
assert result["total_models"] == 500
|
||||||
|
assert result["pending_count"] == 300
|
||||||
|
assert result["needs_download"] is True
|
||||||
|
# Per-folder path should run once per candidate.
|
||||||
|
assert per_model_checks == 500
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
@pytest.mark.usefixtures("tmp_path")
|
||||||
|
async def test_check_pending_models_bulk_index_includes_legacy_folders(
|
||||||
|
monkeypatch: pytest.MonkeyPatch,
|
||||||
|
tmp_path,
|
||||||
|
settings_manager,
|
||||||
|
):
|
||||||
|
"""In multi-library mode the bulk index also scans the legacy root so models
|
||||||
|
whose folders have not been consolidated yet are not reported pending."""
|
||||||
|
|
||||||
|
ws_manager = RecordingWebSocketManager()
|
||||||
|
manager = download_module.DownloadManager(ws_manager=ws_manager)
|
||||||
|
|
||||||
|
monkeypatch.setitem(settings_manager.settings, "example_images_path", str(tmp_path))
|
||||||
|
monkeypatch.setitem(settings_manager.settings, "libraries", {"default": {}, "extra": {}})
|
||||||
|
monkeypatch.setitem(settings_manager.settings, "active_library", "extra")
|
||||||
|
|
||||||
|
# 1500 unprocessed models triggers the bulk lookup path
|
||||||
|
models = [
|
||||||
|
{"sha256": f"{i:064x}", "model_name": f"Model {i}"}
|
||||||
|
for i in range(1500)
|
||||||
|
]
|
||||||
|
|
||||||
|
# Folders live at the LEGACY root/<hash> path (not yet consolidated)
|
||||||
|
for i in range(500):
|
||||||
|
model_dir = tmp_path / f"{i:064x}"
|
||||||
|
model_dir.mkdir()
|
||||||
|
(model_dir / "image_0.png").write_text("data")
|
||||||
|
|
||||||
|
_patch_scanners(monkeypatch, lora_scanner=StubScanner(models))
|
||||||
|
|
||||||
|
result = await manager.check_pending_models(["lora"])
|
||||||
|
|
||||||
|
assert result["success"] is True
|
||||||
|
assert result["total_models"] == 1500
|
||||||
|
assert result["pending_count"] == 1000
|
||||||
|
assert result["needs_download"] is True
|
||||||
|
|
||||||
|
|
||||||
@pytest.fixture
|
@pytest.fixture
|
||||||
def settings_manager():
|
def settings_manager():
|
||||||
return get_settings_manager()
|
return get_settings_manager()
|
||||||
|
|||||||
@@ -35,9 +35,11 @@ class DummyDownloader:
|
|||||||
def reset_singletons():
|
def reset_singletons():
|
||||||
CivitaiClient._instance = None
|
CivitaiClient._instance = None
|
||||||
ModelMetadataProviderManager._instance = None
|
ModelMetadataProviderManager._instance = None
|
||||||
|
civitai_client_module._creator_model_count_cache.clear()
|
||||||
yield
|
yield
|
||||||
CivitaiClient._instance = None
|
CivitaiClient._instance = None
|
||||||
ModelMetadataProviderManager._instance = None
|
ModelMetadataProviderManager._instance = None
|
||||||
|
civitai_client_module._creator_model_count_cache.clear()
|
||||||
|
|
||||||
|
|
||||||
@pytest.fixture
|
@pytest.fixture
|
||||||
@@ -622,3 +624,162 @@ async def test_get_image_info_handles_invalid_id(monkeypatch, downloader, caplog
|
|||||||
|
|
||||||
assert result is None
|
assert result is None
|
||||||
assert "Invalid image ID format" in caplog.text
|
assert "Invalid image ID format" in caplog.text
|
||||||
|
|
||||||
|
|
||||||
|
async def test_get_user_models_requests_first_page_with_stable_params(downloader):
|
||||||
|
request_calls = []
|
||||||
|
|
||||||
|
async def fake_make_request(method, url, use_auth=True, **kwargs):
|
||||||
|
request_calls.append({"method": method, "url": url, "kwargs": kwargs})
|
||||||
|
return True, {
|
||||||
|
"items": [
|
||||||
|
{
|
||||||
|
"id": 1,
|
||||||
|
"modelVersions": [
|
||||||
|
{"id": 100, "images": [{"meta": {"comfy": {"x": 1}}}]}
|
||||||
|
],
|
||||||
|
}
|
||||||
|
],
|
||||||
|
"metadata": {"nextCursor": "next-token"},
|
||||||
|
}
|
||||||
|
|
||||||
|
downloader.make_request = fake_make_request
|
||||||
|
|
||||||
|
client = await CivitaiClient.get_instance()
|
||||||
|
result = await client.get_user_models("pixel")
|
||||||
|
|
||||||
|
assert result is not None
|
||||||
|
assert result["nextCursor"] == "next-token"
|
||||||
|
assert len(result["items"]) == 1
|
||||||
|
# comfy metadata is still stripped
|
||||||
|
assert "comfy" not in result["items"][0]["modelVersions"][0]["images"][0]["meta"]
|
||||||
|
|
||||||
|
call = request_calls[0]
|
||||||
|
assert call["method"] == "GET"
|
||||||
|
assert call["url"] == "https://civitai.red/api/v1/models"
|
||||||
|
params = call["kwargs"]["params"]
|
||||||
|
assert params["username"] == "pixel"
|
||||||
|
assert params["nsfw"] == "true"
|
||||||
|
assert params["limit"] == 100
|
||||||
|
assert params["sort"] == "Newest"
|
||||||
|
assert params["period"] == "AllTime"
|
||||||
|
assert "cursor" not in params
|
||||||
|
|
||||||
|
|
||||||
|
async def test_get_user_models_passes_cursor_and_stringifies_next_cursor(downloader):
|
||||||
|
request_calls = []
|
||||||
|
|
||||||
|
async def fake_make_request(method, url, use_auth=True, **kwargs):
|
||||||
|
request_calls.append(kwargs)
|
||||||
|
return True, {"items": [], "metadata": {"nextCursor": 12345}}
|
||||||
|
|
||||||
|
downloader.make_request = fake_make_request
|
||||||
|
|
||||||
|
client = await CivitaiClient.get_instance()
|
||||||
|
result = await client.get_user_models("pixel", cursor="opaque-token")
|
||||||
|
|
||||||
|
assert request_calls[0]["params"]["cursor"] == "opaque-token"
|
||||||
|
assert result == {"items": [], "nextCursor": "12345"}
|
||||||
|
|
||||||
|
|
||||||
|
async def test_get_user_models_without_next_cursor_returns_none_cursor(downloader):
|
||||||
|
async def fake_make_request(method, url, use_auth=True, **kwargs):
|
||||||
|
return True, {"items": [{"id": 1, "modelVersions": []}], "metadata": {}}
|
||||||
|
|
||||||
|
downloader.make_request = fake_make_request
|
||||||
|
|
||||||
|
client = await CivitaiClient.get_instance()
|
||||||
|
result = await client.get_user_models("pixel")
|
||||||
|
|
||||||
|
assert result == {"items": [{"id": 1, "modelVersions": []}], "nextCursor": None}
|
||||||
|
|
||||||
|
|
||||||
|
async def test_get_user_models_failure_returns_none(downloader):
|
||||||
|
async def fake_make_request(method, url, use_auth=True, **kwargs):
|
||||||
|
return False, "500 server error"
|
||||||
|
|
||||||
|
downloader.make_request = fake_make_request
|
||||||
|
|
||||||
|
client = await CivitaiClient.get_instance()
|
||||||
|
result = await client.get_user_models("pixel")
|
||||||
|
|
||||||
|
assert result is None
|
||||||
|
|
||||||
|
|
||||||
|
async def test_get_creator_model_count_matches_exact_username(downloader):
|
||||||
|
request_calls = []
|
||||||
|
|
||||||
|
async def fake_make_request(method, url, use_auth=True, **kwargs):
|
||||||
|
request_calls.append({"url": url, "kwargs": kwargs})
|
||||||
|
return True, {
|
||||||
|
"items": [
|
||||||
|
{"username": "pixelart", "modelCount": 5},
|
||||||
|
{"username": "Pixel", "modelCount": 2140},
|
||||||
|
]
|
||||||
|
}
|
||||||
|
|
||||||
|
downloader.make_request = fake_make_request
|
||||||
|
|
||||||
|
client = await CivitaiClient.get_instance()
|
||||||
|
count = await client.get_creator_model_count("pixel")
|
||||||
|
|
||||||
|
assert count == 2140
|
||||||
|
assert request_calls[0]["url"] == "https://civitai.red/api/v1/creators"
|
||||||
|
assert request_calls[0]["kwargs"]["params"] == {"query": "pixel", "limit": 10}
|
||||||
|
|
||||||
|
|
||||||
|
async def test_get_creator_model_count_without_exact_match_returns_none(downloader):
|
||||||
|
async def fake_make_request(method, url, use_auth=True, **kwargs):
|
||||||
|
return True, {"items": [{"username": "pixelart", "modelCount": 5}]}
|
||||||
|
|
||||||
|
downloader.make_request = fake_make_request
|
||||||
|
|
||||||
|
client = await CivitaiClient.get_instance()
|
||||||
|
count = await client.get_creator_model_count("pixel")
|
||||||
|
|
||||||
|
assert count is None
|
||||||
|
|
||||||
|
|
||||||
|
async def test_get_creator_model_count_caches_results(downloader):
|
||||||
|
request_count = 0
|
||||||
|
|
||||||
|
async def fake_make_request(method, url, use_auth=True, **kwargs):
|
||||||
|
nonlocal request_count
|
||||||
|
request_count += 1
|
||||||
|
return True, {"items": [{"username": "pixel", "modelCount": 42}]}
|
||||||
|
|
||||||
|
downloader.make_request = fake_make_request
|
||||||
|
|
||||||
|
client = await CivitaiClient.get_instance()
|
||||||
|
|
||||||
|
assert await client.get_creator_model_count("pixel") == 42
|
||||||
|
# case-insensitive cache key, second call served from cache
|
||||||
|
assert await client.get_creator_model_count("Pixel") == 42
|
||||||
|
assert request_count == 1
|
||||||
|
|
||||||
|
|
||||||
|
async def test_get_creator_model_count_caches_failures(downloader):
|
||||||
|
request_count = 0
|
||||||
|
|
||||||
|
async def fake_make_request(method, url, use_auth=True, **kwargs):
|
||||||
|
nonlocal request_count
|
||||||
|
request_count += 1
|
||||||
|
return False, "500 server error"
|
||||||
|
|
||||||
|
downloader.make_request = fake_make_request
|
||||||
|
|
||||||
|
client = await CivitaiClient.get_instance()
|
||||||
|
|
||||||
|
assert await client.get_creator_model_count("pixel") is None
|
||||||
|
assert await client.get_creator_model_count("pixel") is None
|
||||||
|
assert request_count == 1
|
||||||
|
|
||||||
|
|
||||||
|
async def test_get_creator_model_count_never_raises(downloader):
|
||||||
|
async def fake_make_request(method, url, use_auth=True, **kwargs):
|
||||||
|
return True, "unexpected non-dict payload"
|
||||||
|
|
||||||
|
downloader.make_request = fake_make_request
|
||||||
|
|
||||||
|
client = await CivitaiClient.get_instance()
|
||||||
|
assert await client.get_creator_model_count("pixel") is None
|
||||||
|
|||||||
@@ -26,6 +26,7 @@ class StubScanner:
|
|||||||
|
|
||||||
def __init__(self, models: list[dict]) -> None:
|
def __init__(self, models: list[dict]) -> None:
|
||||||
self._cache = SimpleNamespace(raw_data=models)
|
self._cache = SimpleNamespace(raw_data=models)
|
||||||
|
self.sync_calls: list[tuple[str, dict]] = []
|
||||||
|
|
||||||
async def get_cached_data(self):
|
async def get_cached_data(self):
|
||||||
return self._cache
|
return self._cache
|
||||||
@@ -38,6 +39,14 @@ class StubScanner:
|
|||||||
break
|
break
|
||||||
return True
|
return True
|
||||||
|
|
||||||
|
async def sync_cache_from_metadata(self, file_path: str, metadata: dict) -> bool:
|
||||||
|
self.sync_calls.append((file_path, metadata))
|
||||||
|
for index, model in enumerate(self._cache.raw_data):
|
||||||
|
if model.get("file_path") == metadata.get("file_path"):
|
||||||
|
self._cache.raw_data[index] = metadata
|
||||||
|
break
|
||||||
|
return True
|
||||||
|
|
||||||
|
|
||||||
def _patch_scanner(monkeypatch: pytest.MonkeyPatch, scanner: StubScanner) -> None:
|
def _patch_scanner(monkeypatch: pytest.MonkeyPatch, scanner: StubScanner) -> None:
|
||||||
async def _get_lora_scanner(cls):
|
async def _get_lora_scanner(cls):
|
||||||
@@ -520,7 +529,8 @@ async def test_not_found_example_images_are_cleaned(
|
|||||||
|
|
||||||
model_dir = images_root / model_hash
|
model_dir = images_root / model_hash
|
||||||
model_dir.mkdir(parents=True, exist_ok=True)
|
model_dir.mkdir(parents=True, exist_ok=True)
|
||||||
(model_dir / "image_0.png").write_bytes(b"first")
|
# Pre-existing file collides with the valid image index (1) so the
|
||||||
|
# pre-download existence check must skip it without a network request
|
||||||
(model_dir / "image_1.png").write_bytes(b"second")
|
(model_dir / "image_1.png").write_bytes(b"second")
|
||||||
|
|
||||||
async def fake_process_local_examples(*_args, **_kwargs):
|
async def fake_process_local_examples(*_args, **_kwargs):
|
||||||
@@ -588,6 +598,9 @@ async def test_not_found_example_images_are_cleaned(
|
|||||||
assert missing_url in downloader.calls
|
assert missing_url in downloader.calls
|
||||||
assert manager._progress["failed_models"] == {model_hash}
|
assert manager._progress["failed_models"] == {model_hash}
|
||||||
assert model_hash in manager._progress["processed_models"]
|
assert model_hash in manager._progress["processed_models"]
|
||||||
|
assert scanner.sync_calls
|
||||||
|
assert len(scanner.sync_calls) == 1
|
||||||
|
assert scanner.sync_calls[0][0] == str(model_path)
|
||||||
|
|
||||||
remaining_images = model_metadata["civitai"]["images"]
|
remaining_images = model_metadata["civitai"]["images"]
|
||||||
assert remaining_images == [
|
assert remaining_images == [
|
||||||
@@ -596,11 +609,188 @@ async def test_not_found_example_images_are_cleaned(
|
|||||||
]
|
]
|
||||||
|
|
||||||
files = sorted(p.name for p in model_dir.iterdir())
|
files = sorted(p.name for p in model_dir.iterdir())
|
||||||
assert files == ["image_0.png", "image_1.png"]
|
assert files == ["image_1.png"]
|
||||||
assert (model_dir / "image_0.png").read_bytes() == b"first"
|
|
||||||
assert (model_dir / "image_1.png").read_bytes() == b"second"
|
assert (model_dir / "image_1.png").read_bytes() == b"second"
|
||||||
|
|
||||||
|
|
||||||
|
async def test_failed_models_retried_when_explicitly_targeted(
|
||||||
|
monkeypatch: pytest.MonkeyPatch,
|
||||||
|
tmp_path,
|
||||||
|
settings_manager,
|
||||||
|
):
|
||||||
|
ws_manager = RecordingWebSocketManager()
|
||||||
|
manager = download_module.DownloadManager(ws_manager=ws_manager)
|
||||||
|
|
||||||
|
images_root = tmp_path / "examples"
|
||||||
|
monkeypatch.setitem(settings_manager.settings, "example_images_path", str(images_root))
|
||||||
|
|
||||||
|
model_hash = "a" * 64
|
||||||
|
model_path = tmp_path / "model.safetensors"
|
||||||
|
model_path.write_text("data", encoding="utf-8")
|
||||||
|
|
||||||
|
model_metadata = {
|
||||||
|
"sha256": model_hash,
|
||||||
|
"model_name": "Failed Example",
|
||||||
|
"file_path": str(model_path),
|
||||||
|
"file_name": "model.safetensors",
|
||||||
|
"civitai": {"images": [{"url": "https://example.com/valid.png"}]},
|
||||||
|
}
|
||||||
|
|
||||||
|
scanner = StubScanner([model_metadata.copy()])
|
||||||
|
_patch_scanner(monkeypatch, scanner)
|
||||||
|
|
||||||
|
# Persist a previous failure so the skip path is exercised
|
||||||
|
images_root.mkdir(parents=True, exist_ok=True)
|
||||||
|
(images_root / ".download_progress.json").write_text(
|
||||||
|
json.dumps(
|
||||||
|
{
|
||||||
|
"failed_models": [model_hash],
|
||||||
|
"processed_models": [],
|
||||||
|
"rate_limited_models": [],
|
||||||
|
}
|
||||||
|
),
|
||||||
|
encoding="utf-8",
|
||||||
|
)
|
||||||
|
|
||||||
|
async def fake_process_local_examples(*_args, **_kwargs):
|
||||||
|
return False
|
||||||
|
|
||||||
|
async def fake_get_updated_model(model_hash_arg, _scanner):
|
||||||
|
return model_metadata
|
||||||
|
|
||||||
|
class DownloaderStub:
|
||||||
|
def __init__(self):
|
||||||
|
self.calls: list[str] = []
|
||||||
|
|
||||||
|
async def download_to_memory(self, url, *_args, **_kwargs):
|
||||||
|
self.calls.append(url)
|
||||||
|
return True, b"\x89PNG\r\n\x1a\n", {"content-type": "image/png"}
|
||||||
|
|
||||||
|
downloader = DownloaderStub()
|
||||||
|
|
||||||
|
async def fake_get_downloader():
|
||||||
|
return downloader
|
||||||
|
|
||||||
|
monkeypatch.setattr(
|
||||||
|
download_module.ExampleImagesProcessor,
|
||||||
|
"process_local_examples",
|
||||||
|
staticmethod(fake_process_local_examples),
|
||||||
|
)
|
||||||
|
monkeypatch.setattr(
|
||||||
|
download_module.MetadataUpdater,
|
||||||
|
"get_updated_model",
|
||||||
|
staticmethod(fake_get_updated_model),
|
||||||
|
)
|
||||||
|
monkeypatch.setattr(download_module, "get_downloader", fake_get_downloader)
|
||||||
|
|
||||||
|
# Without explicit hashes the previously failed model is skipped
|
||||||
|
skipped_manager = download_module.DownloadManager(ws_manager=RecordingWebSocketManager())
|
||||||
|
result = await skipped_manager.start_download({"model_types": ["lora"], "delay": 0})
|
||||||
|
assert result["success"] is True
|
||||||
|
if skipped_manager._download_task is not None:
|
||||||
|
await asyncio.wait_for(skipped_manager._download_task, timeout=1)
|
||||||
|
assert downloader.calls == []
|
||||||
|
|
||||||
|
# With explicit hashes the previously failed model is retried and cleared
|
||||||
|
result = await manager.start_download(
|
||||||
|
{"model_types": ["lora"], "delay": 0, "model_hashes": [model_hash]}
|
||||||
|
)
|
||||||
|
assert result["success"] is True
|
||||||
|
if manager._download_task is not None:
|
||||||
|
await asyncio.wait_for(manager._download_task, timeout=1)
|
||||||
|
assert downloader.calls == ["https://example.com/valid.png"]
|
||||||
|
assert manager._progress["failed_models"] == set()
|
||||||
|
assert model_hash in manager._progress["processed_models"]
|
||||||
|
|
||||||
|
|
||||||
|
async def test_explicit_targets_fill_partial_example_gaps(
|
||||||
|
monkeypatch: pytest.MonkeyPatch,
|
||||||
|
tmp_path,
|
||||||
|
settings_manager,
|
||||||
|
):
|
||||||
|
ws_manager = RecordingWebSocketManager()
|
||||||
|
|
||||||
|
images_root = tmp_path / "examples"
|
||||||
|
monkeypatch.setitem(settings_manager.settings, "example_images_path", str(images_root))
|
||||||
|
|
||||||
|
model_hash = "b" * 64
|
||||||
|
model_path = tmp_path / "model.safetensors"
|
||||||
|
model_path.write_text("data", encoding="utf-8")
|
||||||
|
|
||||||
|
model_metadata = {
|
||||||
|
"sha256": model_hash,
|
||||||
|
"model_name": "Partial Example",
|
||||||
|
"file_path": str(model_path),
|
||||||
|
"file_name": "model.safetensors",
|
||||||
|
"civitai": {
|
||||||
|
"images": [
|
||||||
|
{"url": "https://example.com/first.png"},
|
||||||
|
{"url": "https://example.com/second.png"},
|
||||||
|
]
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
scanner = StubScanner([model_metadata.copy()])
|
||||||
|
_patch_scanner(monkeypatch, scanner)
|
||||||
|
|
||||||
|
# Simulate a partially populated folder: index 0 already downloaded
|
||||||
|
model_dir = images_root / model_hash
|
||||||
|
model_dir.mkdir(parents=True, exist_ok=True)
|
||||||
|
(model_dir / "image_0.png").write_bytes(b"existing")
|
||||||
|
|
||||||
|
async def fake_process_local_examples(*_args, **_kwargs):
|
||||||
|
return False
|
||||||
|
|
||||||
|
async def fake_get_updated_model(model_hash_arg, _scanner):
|
||||||
|
return model_metadata
|
||||||
|
|
||||||
|
class DownloaderStub:
|
||||||
|
def __init__(self):
|
||||||
|
self.calls: list[str] = []
|
||||||
|
|
||||||
|
async def download_to_memory(self, url, *_args, **_kwargs):
|
||||||
|
self.calls.append(url)
|
||||||
|
return True, b"\x89PNG\r\n\x1a\n", {"content-type": "image/png"}
|
||||||
|
|
||||||
|
downloader = DownloaderStub()
|
||||||
|
|
||||||
|
async def fake_get_downloader():
|
||||||
|
return downloader
|
||||||
|
|
||||||
|
monkeypatch.setattr(
|
||||||
|
download_module.ExampleImagesProcessor,
|
||||||
|
"process_local_examples",
|
||||||
|
staticmethod(fake_process_local_examples),
|
||||||
|
)
|
||||||
|
monkeypatch.setattr(
|
||||||
|
download_module.MetadataUpdater,
|
||||||
|
"get_updated_model",
|
||||||
|
staticmethod(fake_get_updated_model),
|
||||||
|
)
|
||||||
|
monkeypatch.setattr(download_module, "get_downloader", fake_get_downloader)
|
||||||
|
|
||||||
|
# Untargeted run treats the populated folder as done
|
||||||
|
untargeted = download_module.DownloadManager(ws_manager=RecordingWebSocketManager())
|
||||||
|
result = await untargeted.start_download({"model_types": ["lora"], "delay": 0})
|
||||||
|
assert result["success"] is True
|
||||||
|
if untargeted._download_task is not None:
|
||||||
|
await asyncio.wait_for(untargeted._download_task, timeout=1)
|
||||||
|
assert downloader.calls == []
|
||||||
|
|
||||||
|
# Explicitly targeted run fills only the missing index, skipping the
|
||||||
|
# existing file without a network request
|
||||||
|
targeted = download_module.DownloadManager(ws_manager=ws_manager)
|
||||||
|
result = await targeted.start_download(
|
||||||
|
{"model_types": ["lora"], "delay": 0, "model_hashes": [model_hash]}
|
||||||
|
)
|
||||||
|
assert result["success"] is True
|
||||||
|
if targeted._download_task is not None:
|
||||||
|
await asyncio.wait_for(targeted._download_task, timeout=1)
|
||||||
|
assert downloader.calls == ["https://example.com/second.png"]
|
||||||
|
assert (model_dir / "image_1.png").exists()
|
||||||
|
assert (model_dir / "image_0.png").read_bytes() == b"existing"
|
||||||
|
|
||||||
|
|
||||||
@pytest.fixture
|
@pytest.fixture
|
||||||
def settings_manager():
|
def settings_manager():
|
||||||
return get_settings_manager()
|
return get_settings_manager()
|
||||||
|
|||||||
@@ -884,7 +884,7 @@ async def test_sync_cache_conditional_resort_skipped(tmp_path: Path, monkeypatch
|
|||||||
raw_data=[dict(entry)], folders=[], name_display_mode="model_name"
|
raw_data=[dict(entry)], folders=[], name_display_mode="model_name"
|
||||||
)
|
)
|
||||||
await scanner._cache.resort()
|
await scanner._cache.resort()
|
||||||
scanner._cache._last_sort = ("name", "asc") # name sort is active
|
scanner._cache._last_sort = ("name", "asc", None) # name sort is active
|
||||||
scanner._tags_count = {"alpha": 1}
|
scanner._tags_count = {"alpha": 1}
|
||||||
scanner._hash_index.add_entry("abc123", "/m/a.safetensors")
|
scanner._hash_index.add_entry("abc123", "/m/a.safetensors")
|
||||||
|
|
||||||
@@ -935,7 +935,7 @@ async def test_sync_cache_conditional_resort_triggered(tmp_path: Path, monkeypat
|
|||||||
raw_data=[dict(entry)], folders=[], name_display_mode="model_name"
|
raw_data=[dict(entry)], folders=[], name_display_mode="model_name"
|
||||||
)
|
)
|
||||||
await scanner._cache.resort()
|
await scanner._cache.resort()
|
||||||
scanner._cache._last_sort = ("name", "asc")
|
scanner._cache._last_sort = ("name", "asc", None)
|
||||||
scanner._tags_count = {"alpha": 1}
|
scanner._tags_count = {"alpha": 1}
|
||||||
scanner._hash_index.add_entry("abc123", "/m/a.safetensors")
|
scanner._hash_index.add_entry("abc123", "/m/a.safetensors")
|
||||||
|
|
||||||
|
|||||||
@@ -0,0 +1,97 @@
|
|||||||
|
"""Tests for sort parsing and the seeded random sort mode."""
|
||||||
|
|
||||||
|
import asyncio
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
|
||||||
|
from py.services.model_cache import ModelCache
|
||||||
|
from py.services.model_query import ModelCacheRepository, SortParams
|
||||||
|
|
||||||
|
|
||||||
|
def _make_cache(items):
|
||||||
|
return ModelCache(
|
||||||
|
raw_data=[
|
||||||
|
{
|
||||||
|
"file_path": f"/models/{name}.safetensors",
|
||||||
|
"file_name": f"{name}.safetensors",
|
||||||
|
"model_name": name,
|
||||||
|
"folder": "",
|
||||||
|
"size": 100,
|
||||||
|
"modified": 0.0,
|
||||||
|
}
|
||||||
|
for name in items
|
||||||
|
],
|
||||||
|
folders=[],
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class TestParseSort:
|
||||||
|
def test_random_with_seed(self):
|
||||||
|
params = ModelCacheRepository.parse_sort("random:abc123")
|
||||||
|
assert params == SortParams(key="random", order="asc", seed="abc123")
|
||||||
|
|
||||||
|
def test_random_without_seed(self):
|
||||||
|
params = ModelCacheRepository.parse_sort("random")
|
||||||
|
assert params == SortParams(key="random", order="asc", seed=None)
|
||||||
|
|
||||||
|
def test_random_empty_seed_falls_back_to_none(self):
|
||||||
|
params = ModelCacheRepository.parse_sort("random:")
|
||||||
|
assert params.seed is None
|
||||||
|
|
||||||
|
def test_regular_sorts_unaffected(self):
|
||||||
|
params = ModelCacheRepository.parse_sort("name:desc")
|
||||||
|
assert params == SortParams(key="name", order="desc", seed=None)
|
||||||
|
|
||||||
|
|
||||||
|
class TestRandomShuffle:
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_same_seed_yields_same_order(self):
|
||||||
|
cache = _make_cache(["a", "b", "c", "d", "e"])
|
||||||
|
await asyncio.sleep(0) # allow background resort task to run
|
||||||
|
|
||||||
|
first = await cache.get_sorted_data("random", "asc", "seed1")
|
||||||
|
second = await cache.get_sorted_data("random", "asc", "seed1")
|
||||||
|
|
||||||
|
assert [item["model_name"] for item in first] == [
|
||||||
|
item["model_name"] for item in second
|
||||||
|
]
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_different_seeds_yield_different_orders(self):
|
||||||
|
cache = _make_cache([f"m{i}" for i in range(20)])
|
||||||
|
await asyncio.sleep(0)
|
||||||
|
|
||||||
|
first = await cache.get_sorted_data("random", "asc", "seed-a")
|
||||||
|
second = await cache.get_sorted_data("random", "asc", "seed-b")
|
||||||
|
|
||||||
|
assert [item["model_name"] for item in first] != [
|
||||||
|
item["model_name"] for item in second
|
||||||
|
]
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_shuffle_is_a_permutation(self):
|
||||||
|
cache = _make_cache(["a", "b", "c", "d", "e"])
|
||||||
|
await asyncio.sleep(0)
|
||||||
|
|
||||||
|
shuffled = await cache.get_sorted_data("random", "asc", "seed")
|
||||||
|
|
||||||
|
assert sorted(item["model_name"] for item in shuffled) == [
|
||||||
|
"a",
|
||||||
|
"b",
|
||||||
|
"c",
|
||||||
|
"d",
|
||||||
|
"e",
|
||||||
|
]
|
||||||
|
assert len({item["file_path"] for item in shuffled}) == 5
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_missing_seed_is_stable(self):
|
||||||
|
cache = _make_cache(["a", "b", "c", "d", "e"])
|
||||||
|
await asyncio.sleep(0)
|
||||||
|
|
||||||
|
first = await cache.get_sorted_data("random", "asc")
|
||||||
|
second = await cache.get_sorted_data("random", "asc")
|
||||||
|
|
||||||
|
assert [item["model_name"] for item in first] == [
|
||||||
|
item["model_name"] for item in second
|
||||||
|
]
|
||||||
@@ -27,6 +27,13 @@ class FakeScanner:
|
|||||||
return list(self._roots)
|
return list(self._roots)
|
||||||
|
|
||||||
|
|
||||||
|
class StubSettings:
|
||||||
|
"""Settings stub that returns defaults, avoiding the real settings singleton."""
|
||||||
|
|
||||||
|
def get(self, key, default=None):
|
||||||
|
return default
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_search_relative_paths_supports_multiple_tokens():
|
async def test_search_relative_paths_supports_multiple_tokens():
|
||||||
scanner = FakeScanner(
|
scanner = FakeScanner(
|
||||||
@@ -101,3 +108,274 @@ async def test_search_safe_does_not_match_all_files():
|
|||||||
matching = await service.search_relative_paths("safe")
|
matching = await service.search_relative_paths("safe")
|
||||||
|
|
||||||
assert len(matching) == 0
|
assert len(matching) == 0
|
||||||
|
|
||||||
|
|
||||||
|
class SfwStubSettings(StubSettings):
|
||||||
|
"""Settings stub with the global SFW filter enabled."""
|
||||||
|
|
||||||
|
def get(self, key, default=None):
|
||||||
|
if key == "show_only_sfw":
|
||||||
|
return True
|
||||||
|
return default
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_search_relative_paths_respects_global_sfw_setting():
|
||||||
|
"""Filtered search applies show_only_sfw like the list endpoint (parity)."""
|
||||||
|
scanner = FakeScanner(
|
||||||
|
[
|
||||||
|
{"file_path": "/models/sfw-model.safetensors", "preview_nsfw_level": 0},
|
||||||
|
{"file_path": "/models/nsfw-model.safetensors", "preview_nsfw_level": 4},
|
||||||
|
],
|
||||||
|
["/models"],
|
||||||
|
)
|
||||||
|
service = DummyService(
|
||||||
|
"stub", scanner, BaseModelMetadata, settings_provider=SfwStubSettings()
|
||||||
|
)
|
||||||
|
|
||||||
|
matching = await service.search_relative_paths("model", apply_filters=True)
|
||||||
|
|
||||||
|
assert matching == ["sfw-model.safetensors"]
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_search_relative_paths_sfw_only_applied_when_filter_mode_is_on():
|
||||||
|
"""Global settings (show_only_sfw) apply only when the filter pipeline runs."""
|
||||||
|
scanner = FakeScanner(
|
||||||
|
[
|
||||||
|
{"file_path": "/models/sfw-model.safetensors", "preview_nsfw_level": 0},
|
||||||
|
{"file_path": "/models/nsfw-model.safetensors", "preview_nsfw_level": 4},
|
||||||
|
],
|
||||||
|
["/models"],
|
||||||
|
)
|
||||||
|
service = DummyService(
|
||||||
|
"stub", scanner, BaseModelMetadata, settings_provider=SfwStubSettings()
|
||||||
|
)
|
||||||
|
|
||||||
|
default_matching = await service.search_relative_paths("model")
|
||||||
|
|
||||||
|
assert default_matching == [
|
||||||
|
"sfw-model.safetensors",
|
||||||
|
"nsfw-model.safetensors",
|
||||||
|
]
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_search_relative_paths_folder_filter_recursive():
|
||||||
|
"""folder filter with recursive=True (default) matches subfolders."""
|
||||||
|
scanner = FakeScanner(
|
||||||
|
[
|
||||||
|
{"file_path": "/models/anime/model-a.safetensors", "folder": "anime"},
|
||||||
|
{
|
||||||
|
"file_path": "/models/anime/nsfw/model-b.safetensors",
|
||||||
|
"folder": "anime/nsfw",
|
||||||
|
},
|
||||||
|
{"file_path": "/models/realistic/model-c.safetensors", "folder": "realistic"},
|
||||||
|
],
|
||||||
|
["/models"],
|
||||||
|
)
|
||||||
|
service = DummyService(
|
||||||
|
"stub", scanner, BaseModelMetadata, settings_provider=StubSettings()
|
||||||
|
)
|
||||||
|
|
||||||
|
matching = await service.search_relative_paths("model", folder="anime")
|
||||||
|
|
||||||
|
assert matching == [
|
||||||
|
f"anime{os.sep}model-a.safetensors",
|
||||||
|
f"anime{os.sep}nsfw{os.sep}model-b.safetensors",
|
||||||
|
]
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_search_relative_paths_folder_filter_exact():
|
||||||
|
"""folder filter with recursive=False matches only the exact folder."""
|
||||||
|
scanner = FakeScanner(
|
||||||
|
[
|
||||||
|
{"file_path": "/models/anime/model-a.safetensors", "folder": "anime"},
|
||||||
|
{
|
||||||
|
"file_path": "/models/anime/nsfw/model-b.safetensors",
|
||||||
|
"folder": "anime/nsfw",
|
||||||
|
},
|
||||||
|
{"file_path": "/models/realistic/model-c.safetensors", "folder": "realistic"},
|
||||||
|
],
|
||||||
|
["/models"],
|
||||||
|
)
|
||||||
|
service = DummyService(
|
||||||
|
"stub", scanner, BaseModelMetadata, settings_provider=StubSettings()
|
||||||
|
)
|
||||||
|
|
||||||
|
matching = await service.search_relative_paths(
|
||||||
|
"model", folder="anime", recursive=False
|
||||||
|
)
|
||||||
|
|
||||||
|
assert matching == [f"anime{os.sep}model-a.safetensors"]
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_search_relative_paths_base_model_filter():
|
||||||
|
scanner = FakeScanner(
|
||||||
|
[
|
||||||
|
{"file_path": "/models/model-a.safetensors", "base_model": "SD 1.5"},
|
||||||
|
{"file_path": "/models/model-b.safetensors", "base_model": "SDXL"},
|
||||||
|
],
|
||||||
|
["/models"],
|
||||||
|
)
|
||||||
|
service = DummyService(
|
||||||
|
"stub", scanner, BaseModelMetadata, settings_provider=StubSettings()
|
||||||
|
)
|
||||||
|
|
||||||
|
matching = await service.search_relative_paths("model", base_models=["SD 1.5"])
|
||||||
|
|
||||||
|
assert matching == ["model-a.safetensors"]
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_search_relative_paths_tag_include():
|
||||||
|
scanner = FakeScanner(
|
||||||
|
[
|
||||||
|
{"file_path": "/models/model-a.safetensors", "tags": ["anime"]},
|
||||||
|
{"file_path": "/models/model-b.safetensors", "tags": ["realistic"]},
|
||||||
|
{"file_path": "/models/model-c.safetensors", "tags": ["anime", "realistic"]},
|
||||||
|
],
|
||||||
|
["/models"],
|
||||||
|
)
|
||||||
|
service = DummyService(
|
||||||
|
"stub", scanner, BaseModelMetadata, settings_provider=StubSettings()
|
||||||
|
)
|
||||||
|
|
||||||
|
matching = await service.search_relative_paths("model", tags={"anime": "include"})
|
||||||
|
|
||||||
|
assert set(matching) == {"model-a.safetensors", "model-c.safetensors"}
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_search_relative_paths_tag_exclude():
|
||||||
|
scanner = FakeScanner(
|
||||||
|
[
|
||||||
|
{"file_path": "/models/model-a.safetensors", "tags": ["anime"]},
|
||||||
|
{"file_path": "/models/model-b.safetensors", "tags": ["realistic"]},
|
||||||
|
],
|
||||||
|
["/models"],
|
||||||
|
)
|
||||||
|
service = DummyService(
|
||||||
|
"stub", scanner, BaseModelMetadata, settings_provider=StubSettings()
|
||||||
|
)
|
||||||
|
|
||||||
|
matching = await service.search_relative_paths("model", tags={"anime": "exclude"})
|
||||||
|
|
||||||
|
assert matching == ["model-b.safetensors"]
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_search_relative_paths_auto_tag_include():
|
||||||
|
scanner = FakeScanner(
|
||||||
|
[
|
||||||
|
{
|
||||||
|
"file_path": "/models/model-i2v.safetensors",
|
||||||
|
"file_name": "model-i2v.safetensors",
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"file_path": "/models/model-t2v.safetensors",
|
||||||
|
"file_name": "model-t2v.safetensors",
|
||||||
|
},
|
||||||
|
],
|
||||||
|
["/models"],
|
||||||
|
)
|
||||||
|
service = DummyService(
|
||||||
|
"stub", scanner, BaseModelMetadata, settings_provider=StubSettings()
|
||||||
|
)
|
||||||
|
|
||||||
|
matching = await service.search_relative_paths(
|
||||||
|
"model", auto_tags={"I2V": "include"}
|
||||||
|
)
|
||||||
|
|
||||||
|
assert matching == ["model-i2v.safetensors"]
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_search_relative_paths_tag_logic_all():
|
||||||
|
scanner = FakeScanner(
|
||||||
|
[
|
||||||
|
{"file_path": "/models/model-a.safetensors", "tags": ["anime", "style"]},
|
||||||
|
{"file_path": "/models/model-b.safetensors", "tags": ["anime"]},
|
||||||
|
],
|
||||||
|
["/models"],
|
||||||
|
)
|
||||||
|
service = DummyService(
|
||||||
|
"stub", scanner, BaseModelMetadata, settings_provider=StubSettings()
|
||||||
|
)
|
||||||
|
|
||||||
|
matching = await service.search_relative_paths(
|
||||||
|
"model", tags={"anime": "include", "style": "include"}, tag_logic="all"
|
||||||
|
)
|
||||||
|
|
||||||
|
assert matching == ["model-a.safetensors"]
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_search_relative_paths_credit_required_filter():
|
||||||
|
# license_flags bit0: 1 = no credit required, 0 = credit required
|
||||||
|
scanner = FakeScanner(
|
||||||
|
[
|
||||||
|
{"file_path": "/models/model-a.safetensors", "license_flags": 127},
|
||||||
|
{"file_path": "/models/model-b.safetensors", "license_flags": 0},
|
||||||
|
],
|
||||||
|
["/models"],
|
||||||
|
)
|
||||||
|
service = DummyService(
|
||||||
|
"stub", scanner, BaseModelMetadata, settings_provider=StubSettings()
|
||||||
|
)
|
||||||
|
|
||||||
|
matching = await service.search_relative_paths("model", credit_required=True)
|
||||||
|
assert matching == ["model-b.safetensors"]
|
||||||
|
|
||||||
|
matching = await service.search_relative_paths("model", credit_required=False)
|
||||||
|
assert matching == ["model-a.safetensors"]
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_search_relative_paths_allow_selling_filter():
|
||||||
|
# license_flags bit1: 1 = commercial image use allowed, 0 = not allowed
|
||||||
|
scanner = FakeScanner(
|
||||||
|
[
|
||||||
|
{"file_path": "/models/model-a.safetensors", "license_flags": 2},
|
||||||
|
{"file_path": "/models/model-b.safetensors", "license_flags": 1},
|
||||||
|
],
|
||||||
|
["/models"],
|
||||||
|
)
|
||||||
|
service = DummyService(
|
||||||
|
"stub", scanner, BaseModelMetadata, settings_provider=StubSettings()
|
||||||
|
)
|
||||||
|
|
||||||
|
matching = await service.search_relative_paths(
|
||||||
|
"model", allow_selling_generated_content=True
|
||||||
|
)
|
||||||
|
assert matching == ["model-a.safetensors"]
|
||||||
|
|
||||||
|
matching = await service.search_relative_paths(
|
||||||
|
"model", allow_selling_generated_content=False
|
||||||
|
)
|
||||||
|
assert matching == ["model-b.safetensors"]
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_search_relative_paths_no_filters_regression():
|
||||||
|
"""No filter kwargs -> behavior is byte-identical to plain token matching."""
|
||||||
|
scanner = FakeScanner(
|
||||||
|
[
|
||||||
|
{"file_path": "/models/flux/detail-model.safetensors"},
|
||||||
|
{"file_path": "/models/flux/only-flux.safetensors"},
|
||||||
|
],
|
||||||
|
["/models"],
|
||||||
|
)
|
||||||
|
service = DummyService(
|
||||||
|
"stub", scanner, BaseModelMetadata, settings_provider=StubSettings()
|
||||||
|
)
|
||||||
|
|
||||||
|
matching = await service.search_relative_paths("flux")
|
||||||
|
|
||||||
|
assert matching == [
|
||||||
|
f"flux{os.sep}only-flux.safetensors",
|
||||||
|
f"flux{os.sep}detail-model.safetensors",
|
||||||
|
]
|
||||||
|
|
||||||
|
|||||||
@@ -860,6 +860,59 @@ def test_set_recipes_path_rewrites_symlinked_recipe_metadata(manager, tmp_path):
|
|||||||
assert not old_json_path.exists()
|
assert not old_json_path.exists()
|
||||||
|
|
||||||
|
|
||||||
|
def test_set_recipes_path_allows_cross_drive_migration(manager, tmp_path, monkeypatch):
|
||||||
|
# Windows regression: os.path.commonpath raises ValueError for paths on
|
||||||
|
# different drives (ntpath semantics). Cross-drive moves must succeed.
|
||||||
|
lora_root = tmp_path / "loras"
|
||||||
|
old_recipes_dir = lora_root / "recipes" / "nested"
|
||||||
|
old_recipes_dir.mkdir(parents=True)
|
||||||
|
manager.set("folder_paths", {"loras": [str(lora_root)]})
|
||||||
|
|
||||||
|
recipe_id = "recipe-cross-drive"
|
||||||
|
old_image_path = old_recipes_dir / f"{recipe_id}.webp"
|
||||||
|
old_json_path = old_recipes_dir / f"{recipe_id}.recipe.json"
|
||||||
|
old_image_path.write_bytes(b"image-bytes")
|
||||||
|
old_json_path.write_text(
|
||||||
|
json.dumps(
|
||||||
|
{
|
||||||
|
"id": recipe_id,
|
||||||
|
"file_path": str(old_image_path),
|
||||||
|
"title": "Recipe Cross Drive",
|
||||||
|
}
|
||||||
|
),
|
||||||
|
encoding="utf-8",
|
||||||
|
)
|
||||||
|
|
||||||
|
new_recipes_dir = tmp_path / "N_drive" / "AI" / "Library" / "Recipes"
|
||||||
|
|
||||||
|
# The effective current recipes dir (source of the migration) is
|
||||||
|
# lora_root/recipes — the nested subdirectory holds the recipe files.
|
||||||
|
source = str(lora_root / "recipes")
|
||||||
|
target = str(new_recipes_dir)
|
||||||
|
real_commonpath = os.path.commonpath
|
||||||
|
|
||||||
|
def fake_commonpath(paths):
|
||||||
|
# Simulate ntpath on Windows: a source/target pair on different
|
||||||
|
# drives shares no common root and raises ValueError.
|
||||||
|
if {source, target} <= set(paths):
|
||||||
|
raise ValueError("Paths don't have the same drive")
|
||||||
|
return real_commonpath(paths)
|
||||||
|
|
||||||
|
monkeypatch.setattr(os.path, "commonpath", fake_commonpath)
|
||||||
|
|
||||||
|
manager.set("recipes_path", str(new_recipes_dir))
|
||||||
|
|
||||||
|
migrated_image_path = new_recipes_dir / "nested" / f"{recipe_id}.webp"
|
||||||
|
migrated_json_path = new_recipes_dir / "nested" / f"{recipe_id}.recipe.json"
|
||||||
|
|
||||||
|
assert manager.get("recipes_path") == str(new_recipes_dir.resolve())
|
||||||
|
assert migrated_image_path.read_bytes() == b"image-bytes"
|
||||||
|
migrated_payload = json.loads(migrated_json_path.read_text(encoding="utf-8"))
|
||||||
|
assert migrated_payload["file_path"] == str(migrated_image_path)
|
||||||
|
assert not old_image_path.exists()
|
||||||
|
assert not old_json_path.exists()
|
||||||
|
|
||||||
|
|
||||||
def test_set_recipes_path_rejects_file_target(manager, tmp_path):
|
def test_set_recipes_path_rejects_file_target(manager, tmp_path):
|
||||||
lora_root = tmp_path / "loras"
|
lora_root = tmp_path / "loras"
|
||||||
lora_root.mkdir()
|
lora_root.mkdir()
|
||||||
|
|||||||
@@ -63,7 +63,7 @@ async def test_start_download_bootstraps_progress_and_task(
|
|||||||
release = asyncio.Event()
|
release = asyncio.Event()
|
||||||
|
|
||||||
async def fake_download(
|
async def fake_download(
|
||||||
self, output_dir, optimize, model_types, delay, library_name, force=False
|
self, output_dir, optimize, model_types, delay, library_name, force=False, model_hashes=None
|
||||||
):
|
):
|
||||||
started.set()
|
started.set()
|
||||||
await release.wait()
|
await release.wait()
|
||||||
@@ -93,6 +93,44 @@ async def test_start_download_bootstraps_progress_and_task(
|
|||||||
assert manager._progress["status"] == "completed"
|
assert manager._progress["status"] == "completed"
|
||||||
|
|
||||||
|
|
||||||
|
async def test_start_download_forwards_model_hashes(
|
||||||
|
monkeypatch: pytest.MonkeyPatch, tmp_path
|
||||||
|
) -> None:
|
||||||
|
settings_manager = get_settings_manager()
|
||||||
|
settings_manager.settings["example_images_path"] = str(tmp_path)
|
||||||
|
settings_manager.settings["libraries"] = {"default": {}}
|
||||||
|
settings_manager.settings["active_library"] = "default"
|
||||||
|
|
||||||
|
manager = download_module.DownloadManager(ws_manager=RecordingWebSocketManager())
|
||||||
|
|
||||||
|
received: Dict[str, Any] = {}
|
||||||
|
|
||||||
|
async def fake_download(
|
||||||
|
self, output_dir, optimize, model_types, delay, library_name, force=False, model_hashes=None
|
||||||
|
):
|
||||||
|
received["model_hashes"] = model_hashes
|
||||||
|
async with self._state_lock:
|
||||||
|
self._is_downloading = False
|
||||||
|
self._download_task = None
|
||||||
|
self._progress["status"] = "completed"
|
||||||
|
|
||||||
|
monkeypatch.setattr(
|
||||||
|
download_module.DownloadManager,
|
||||||
|
"_download_all_example_images",
|
||||||
|
fake_download,
|
||||||
|
)
|
||||||
|
|
||||||
|
result = await manager.start_download(
|
||||||
|
{"model_types": ["lora"], "delay": 0, "model_hashes": ["abc123", "def456"]}
|
||||||
|
)
|
||||||
|
assert result["success"] is True
|
||||||
|
|
||||||
|
task = manager._download_task
|
||||||
|
assert task is not None
|
||||||
|
await asyncio.wait_for(task, timeout=1)
|
||||||
|
assert received["model_hashes"] == ["abc123", "def456"]
|
||||||
|
|
||||||
|
|
||||||
async def test_pause_and_resume_flow(monkeypatch: pytest.MonkeyPatch, tmp_path) -> None:
|
async def test_pause_and_resume_flow(monkeypatch: pytest.MonkeyPatch, tmp_path) -> None:
|
||||||
settings_manager = get_settings_manager()
|
settings_manager = get_settings_manager()
|
||||||
settings_manager.settings["example_images_path"] = str(tmp_path)
|
settings_manager.settings["example_images_path"] = str(tmp_path)
|
||||||
|
|||||||
@@ -15,6 +15,7 @@ class StubScanner:
|
|||||||
def __init__(self, cache_items: List[Dict[str, Any]]) -> None:
|
def __init__(self, cache_items: List[Dict[str, Any]]) -> None:
|
||||||
self.cache = SimpleNamespace(raw_data=cache_items)
|
self.cache = SimpleNamespace(raw_data=cache_items)
|
||||||
self.updates: List[Tuple[str, str, Dict[str, Any]]] = []
|
self.updates: List[Tuple[str, str, Dict[str, Any]]] = []
|
||||||
|
self.sync_updates: List[Tuple[str, Dict[str, Any]]] = []
|
||||||
|
|
||||||
async def get_cached_data(self):
|
async def get_cached_data(self):
|
||||||
return self.cache
|
return self.cache
|
||||||
@@ -23,6 +24,10 @@ class StubScanner:
|
|||||||
self.updates.append((old_path, new_path, metadata))
|
self.updates.append((old_path, new_path, metadata))
|
||||||
return True
|
return True
|
||||||
|
|
||||||
|
async def sync_cache_from_metadata(self, file_path: str, metadata: Dict[str, Any]) -> bool:
|
||||||
|
self.sync_updates.append((file_path, metadata))
|
||||||
|
return True
|
||||||
|
|
||||||
|
|
||||||
@pytest.fixture(autouse=True)
|
@pytest.fixture(autouse=True)
|
||||||
def patch_metadata_manager(monkeypatch: pytest.MonkeyPatch):
|
def patch_metadata_manager(monkeypatch: pytest.MonkeyPatch):
|
||||||
@@ -83,7 +88,7 @@ async def test_update_metadata_after_import_enriches_entries(monkeypatch: pytest
|
|||||||
assert custom[0]["type"] == "image"
|
assert custom[0]["type"] == "image"
|
||||||
|
|
||||||
assert Path(patch_metadata_manager[0][0]) == model_file
|
assert Path(patch_metadata_manager[0][0]) == model_file
|
||||||
assert scanner.updates
|
assert scanner.sync_updates
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
@@ -151,8 +156,8 @@ async def test_update_metadata_after_import_preserves_existing_metadata(
|
|||||||
assert saved_payload["civitai"]["trainedWords"] == ["foo"]
|
assert saved_payload["civitai"]["trainedWords"] == ["foo"]
|
||||||
assert {entry["id"] for entry in saved_payload["civitai"]["customImages"]} == {"existing-id", "new-id"}
|
assert {entry["id"] for entry in saved_payload["civitai"]["customImages"]} == {"existing-id", "new-id"}
|
||||||
|
|
||||||
assert scanner.updates
|
assert scanner.sync_updates
|
||||||
updated_metadata = scanner.updates[-1][2]
|
updated_metadata = scanner.sync_updates[-1][1]
|
||||||
assert updated_metadata["civitai"]["images"] == existing_payload["civitai"]["images"]
|
assert updated_metadata["civitai"]["images"] == existing_payload["civitai"]["images"]
|
||||||
assert {entry["id"] for entry in updated_metadata["civitai"]["customImages"]} == {"existing-id", "new-id"}
|
assert {entry["id"] for entry in updated_metadata["civitai"]["customImages"]} == {"existing-id", "new-id"}
|
||||||
|
|
||||||
|
|||||||
@@ -100,6 +100,54 @@ def test_get_file_extension_media_type_hint_low_priority() -> None:
|
|||||||
assert ext == ".mp4"
|
assert ext == ".mp4"
|
||||||
|
|
||||||
|
|
||||||
|
def test_example_image_file_exists_checks_plausible_extensions(tmp_path) -> None:
|
||||||
|
proc = processor_module.ExampleImagesProcessor
|
||||||
|
assert proc._example_image_file_exists(str(tmp_path), 0) is False
|
||||||
|
Path(tmp_path, "image_0.webp").write_bytes(b"x")
|
||||||
|
assert proc._example_image_file_exists(str(tmp_path), 0) is True
|
||||||
|
assert proc._example_image_file_exists(str(tmp_path), 1) is False
|
||||||
|
|
||||||
|
|
||||||
|
def test_example_image_file_exists_video_hint_only_checks_video_extensions(tmp_path) -> None:
|
||||||
|
proc = processor_module.ExampleImagesProcessor
|
||||||
|
Path(tmp_path, "image_2.jpg").write_bytes(b"x")
|
||||||
|
# An existing image file must not satisfy a video-hinted lookup
|
||||||
|
assert proc._example_image_file_exists(str(tmp_path), 2, "video") is False
|
||||||
|
Path(tmp_path, "image_2.mp4").write_bytes(b"x")
|
||||||
|
assert proc._example_image_file_exists(str(tmp_path), 2, "video") is True
|
||||||
|
|
||||||
|
|
||||||
|
async def test_download_model_images_with_tracking_skips_existing_files(tmp_path) -> None:
|
||||||
|
proc = processor_module.ExampleImagesProcessor
|
||||||
|
images = [
|
||||||
|
{"url": "https://image.civitai.com/a/b", "type": "image"},
|
||||||
|
{"url": "https://image.civitai.com/c/d", "type": "image"},
|
||||||
|
]
|
||||||
|
Path(tmp_path, "image_0.jpg").write_bytes(b"existing")
|
||||||
|
|
||||||
|
class RecordingDownloader:
|
||||||
|
def __init__(self) -> None:
|
||||||
|
self.calls: list[str] = []
|
||||||
|
|
||||||
|
async def download_to_memory(self, url, use_auth=False, return_headers=False):
|
||||||
|
self.calls.append(url)
|
||||||
|
return True, b"\xff\xd8\xff" + b"data", {}
|
||||||
|
|
||||||
|
downloader = RecordingDownloader()
|
||||||
|
success, is_stale, failed, rate_limited = await proc.download_model_images_with_tracking(
|
||||||
|
"hash", "model", images, str(tmp_path), False, downloader
|
||||||
|
)
|
||||||
|
|
||||||
|
assert success is True
|
||||||
|
assert is_stale is False
|
||||||
|
assert failed == []
|
||||||
|
assert rate_limited == []
|
||||||
|
# Only the missing image is requested; the existing one is skipped without a network call
|
||||||
|
assert len(downloader.calls) == 1
|
||||||
|
assert "c/d" in downloader.calls[0]
|
||||||
|
assert Path(tmp_path, "image_1.jpg").exists()
|
||||||
|
|
||||||
|
|
||||||
class StubScanner:
|
class StubScanner:
|
||||||
def __init__(self, models: list[Dict[str, Any]]) -> None:
|
def __init__(self, models: list[Dict[str, Any]]) -> None:
|
||||||
self._cache = SimpleNamespace(raw_data=models)
|
self._cache = SimpleNamespace(raw_data=models)
|
||||||
|
|||||||
@@ -288,3 +288,58 @@ class TestIsobmffBrotliExtraction:
|
|||||||
# Direct extraction should return None because decompressed size exceeds limit
|
# Direct extraction should return None because decompressed size exceeds limit
|
||||||
result = ExifUtils._extract_isobmff_brotli(str(path))
|
result = ExifUtils._extract_isobmff_brotli(str(path))
|
||||||
assert result is None
|
assert result is None
|
||||||
|
|
||||||
|
|
||||||
|
# --- get_image_dimensions tests ---
|
||||||
|
|
||||||
|
|
||||||
|
def test_get_image_dimensions_returns_actual_size(tmp_path):
|
||||||
|
"""(a) A valid image returns its real (width, height)."""
|
||||||
|
image_path = tmp_path / "preview.png"
|
||||||
|
Image.new("RGB", (64, 32), color="red").save(image_path)
|
||||||
|
|
||||||
|
assert ExifUtils.get_image_dimensions(str(image_path)) == (64, 32)
|
||||||
|
|
||||||
|
|
||||||
|
def test_get_image_dimensions_missing_path_returns_none(tmp_path):
|
||||||
|
"""(b) A nonexistent path returns None without raising."""
|
||||||
|
assert ExifUtils.get_image_dimensions(str(tmp_path / "missing.png")) is None
|
||||||
|
|
||||||
|
|
||||||
|
def test_get_image_dimensions_skips_video_extension_without_pil(tmp_path, monkeypatch):
|
||||||
|
"""(c) A .mp4 path returns None and never invokes PIL."""
|
||||||
|
video_path = tmp_path / "preview.mp4"
|
||||||
|
video_path.write_bytes(b"not really a video")
|
||||||
|
|
||||||
|
def fail_if_called(*args, **kwargs):
|
||||||
|
raise AssertionError("PIL Image.open must not be called for video paths")
|
||||||
|
|
||||||
|
monkeypatch.setattr("py.utils.exif_utils.Image.open", fail_if_called)
|
||||||
|
|
||||||
|
assert ExifUtils.get_image_dimensions(str(video_path)) is None
|
||||||
|
|
||||||
|
|
||||||
|
def test_get_image_dimensions_corrupt_file_returns_none(tmp_path):
|
||||||
|
"""(d) A corrupt file returns None without raising."""
|
||||||
|
image_path = tmp_path / "corrupt.png"
|
||||||
|
image_path.write_bytes(b"\x00\x01\x02\x03 not a real image")
|
||||||
|
|
||||||
|
assert ExifUtils.get_image_dimensions(str(image_path)) is None
|
||||||
|
|
||||||
|
|
||||||
|
def test_get_image_dimensions_cache_key_includes_mtime(tmp_path):
|
||||||
|
"""(e) Replacing a path with a different-size image returns the new size."""
|
||||||
|
image_path = tmp_path / "replaced.png"
|
||||||
|
Image.new("RGB", (64, 32), color="red").save(image_path)
|
||||||
|
assert ExifUtils.get_image_dimensions(str(image_path)) == (64, 32)
|
||||||
|
|
||||||
|
Image.new("RGB", (100, 50), color="blue").save(image_path)
|
||||||
|
assert ExifUtils.get_image_dimensions(str(image_path)) == (100, 50)
|
||||||
|
|
||||||
|
|
||||||
|
def test_get_image_dimensions_skips_unreadable_formats(tmp_path):
|
||||||
|
"""(f) .avif/.jxl paths return None without raising."""
|
||||||
|
for ext in (".avif", ".jxl"):
|
||||||
|
image_path = tmp_path / f"preview{ext}"
|
||||||
|
image_path.write_bytes(b"fake container data")
|
||||||
|
assert ExifUtils.get_image_dimensions(str(image_path)) is None
|
||||||
|
|||||||
@@ -45,7 +45,7 @@ export interface AutocompleteTextWidgetInterface {
|
|||||||
const props = defineProps<{
|
const props = defineProps<{
|
||||||
widget: AutocompleteTextWidgetInterface
|
widget: AutocompleteTextWidgetInterface
|
||||||
node: { id: number }
|
node: { id: number }
|
||||||
modelType?: 'loras' | 'embeddings' | 'custom_words' | 'prompt'
|
modelType?: 'loras' | 'prompt'
|
||||||
placeholder?: string
|
placeholder?: string
|
||||||
showPreview?: boolean
|
showPreview?: boolean
|
||||||
spellcheck?: boolean
|
spellcheck?: boolean
|
||||||
|
|||||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user