feat(backend): add Other model type (VAE/upscaler/text encoder) scanner, service and routes

This commit is contained in:
Will Miao
2026-09-12 11:25:51 +08:00
parent 3070838a42
commit 27da7b3ca3
22 changed files with 1771 additions and 19 deletions
+1 -1
View File
@@ -166,7 +166,7 @@ The system runs in two modes:
### Model Types & Routes ### Model Types & Routes
- API endpoints follow `/loras/*`, `/checkpoints/*`, `/embeddings/*` patterns - API endpoints follow `/loras/*`, `/checkpoints/*`, `/embeddings/*`, `/other/*` patterns
- Route registrars organize endpoints by domain: `ModelRouteRegistrar`, `RecipeRouteRegistrar`, etc. - Route registrars organize endpoints by domain: `ModelRouteRegistrar`, `RecipeRouteRegistrar`, etc.
- Request handlers in `py/routes/handlers/` implement route logic - Request handlers in `py/routes/handlers/` implement route logic
- All routes use aiohttp, return `web.json_response` or `web.Response` - All routes use aiohttp, return `web.json_response` or `web.Response`
+161 -1
View File
@@ -17,6 +17,10 @@ import types as _types
import time import time
from .utils.cache_paths import CacheType, get_cache_file_path, get_legacy_cache_paths from .utils.cache_paths import CacheType, get_cache_file_path, get_legacy_cache_paths
from .utils.constants import (
DEFAULT_OTHER_MODEL_FOLDERS,
OTHER_MODEL_FOLDER_SUBTYPES,
)
from .utils.settings_paths import ( from .utils.settings_paths import (
ensure_settings_file, ensure_settings_file,
get_settings_dir, get_settings_dir,
@@ -172,6 +176,13 @@ class Config:
self.embeddings_roots = None self.embeddings_roots = None
self.base_models_roots = self._init_checkpoint_paths() self.base_models_roots = self._init_checkpoint_paths()
self.embeddings_roots = self._init_embedding_paths() self.embeddings_roots = self._init_embedding_paths()
# Other-model roots (VAE, upscalers, text encoders, ...): flat deduped
# list plus a normalized root -> sub_type map and per-folder_paths-key
# roots for settings persistence.
self.other_roots: Optional[List[str]] = None
self.other_root_subtypes: Dict[str, str] = {}
self.other_folder_roots: Dict[str, List[str]] = {}
self.other_roots = self._init_other_paths()
# Extra paths (only for LoRA Manager, not shared with ComfyUI) # Extra paths (only for LoRA Manager, not shared with ComfyUI)
self.extra_loras_roots: List[str] = [] self.extra_loras_roots: List[str] = []
self.extra_checkpoints_roots: List[str] = [] self.extra_checkpoints_roots: List[str] = []
@@ -336,6 +347,10 @@ class Config:
"unet": list(self.unet_roots or []), "unet": list(self.unet_roots or []),
"embeddings": list(self.embeddings_roots or []), "embeddings": list(self.embeddings_roots or []),
} }
# Persist the other-model roots under their original folder_paths
# keys so library switching round-trips them.
for key, roots in (self.other_folder_roots or {}).items():
target_folder_paths[key] = list(roots)
normalized_target_paths = _normalize_folder_paths_for_comparison( normalized_target_paths = _normalize_folder_paths_for_comparison(
target_folder_paths target_folder_paths
@@ -522,6 +537,7 @@ class Config:
roots.extend(self.loras_roots or []) roots.extend(self.loras_roots or [])
roots.extend(self.base_models_roots or []) roots.extend(self.base_models_roots or [])
roots.extend(self.embeddings_roots or []) roots.extend(self.embeddings_roots or [])
roots.extend(self.other_roots or [])
# Include extra paths for scanning symlinks # Include extra paths for scanning symlinks
roots.extend(self.extra_loras_roots or []) roots.extend(self.extra_loras_roots or [])
roots.extend(self.extra_checkpoints_roots or []) roots.extend(self.extra_checkpoints_roots or [])
@@ -862,6 +878,8 @@ class Config:
preview_roots.update(self._expand_preview_root(root)) preview_roots.update(self._expand_preview_root(root))
for root in self.embeddings_roots or []: for root in self.embeddings_roots or []:
preview_roots.update(self._expand_preview_root(root)) preview_roots.update(self._expand_preview_root(root))
for root in self.other_roots or []:
preview_roots.update(self._expand_preview_root(root))
# Include extra paths for preview access # Include extra paths for preview access
for root in self.extra_loras_roots or []: for root in self.extra_loras_roots or []:
preview_roots.update(self._expand_preview_root(root)) preview_roots.update(self._expand_preview_root(root))
@@ -882,7 +900,7 @@ class Config:
path for path in preview_roots if path.is_absolute() path for path in preview_roots if path.is_absolute()
} }
logger.debug( logger.debug(
"Preview roots rebuilt: %d paths from %d lora roots (%d extra), %d checkpoint roots (%d extra), %d embedding roots (%d extra), %d symlink mappings", "Preview roots rebuilt: %d paths from %d lora roots (%d extra), %d checkpoint roots (%d extra), %d embedding roots (%d extra), %d other roots, %d symlink mappings",
len(self._preview_root_paths), len(self._preview_root_paths),
len(self.loras_roots or []), len(self.loras_roots or []),
len(self.extra_loras_roots or []), len(self.extra_loras_roots or []),
@@ -890,6 +908,7 @@ class Config:
len(self.extra_checkpoints_roots or []), len(self.extra_checkpoints_roots or []),
len(self.embeddings_roots or []), len(self.embeddings_roots or []),
len(self.extra_embeddings_roots or []), len(self.extra_embeddings_roots or []),
len(self.other_roots or []),
len(self._path_mappings), len(self._path_mappings),
) )
@@ -1128,6 +1147,102 @@ class Config:
return unique_paths return unique_paths
def _get_enabled_other_folder_keys(self) -> List[str]:
"""Return the OTHER_MODEL_FOLDER_SUBTYPES keys that are enabled.
Default-enabled categories come from DEFAULT_OTHER_MODEL_FOLDERS;
opt-in categories (e.g. controlnet) are added via the
``enabled_other_folders`` setting (a list of folder_paths keys).
"""
keys = list(DEFAULT_OTHER_MODEL_FOLDERS)
try:
from .services.settings_manager import get_settings_manager
extra = get_settings_manager().get("enabled_other_folders", [])
except Exception:
extra = []
if isinstance(extra, str):
extra = [extra]
if isinstance(extra, Iterable):
for key in extra:
if (
isinstance(key, str)
and key in OTHER_MODEL_FOLDER_SUBTYPES
and key not in keys
):
keys.append(key)
return keys
def _prepare_other_paths(
self, folder_path_map: Mapping[str, Iterable[str]]
) -> Tuple[List[str], Dict[str, str], Dict[str, List[str]]]:
"""Prepare other-model paths from a folder_paths-key -> raw paths map.
Returns:
Tuple of (all_unique_roots, business_root -> sub_type map,
folder_paths key -> business roots). This method does NOT modify
instance variables - callers must set them.
"""
unique_paths: List[str] = []
sub_type_map: Dict[str, str] = {}
per_key_roots: Dict[str, List[str]] = {}
seen_real_paths: Dict[str, str] = {} # real path -> business path
# Cross-scanner overlap detection: warn when an "other" root is
# already covered by the checkpoints/unet or embeddings scanners.
# Kept (not dropped) on purpose - duplicate cards across pages are
# cosmetic, while dropping would silently unmanage the files.
covered_real_paths = {
os.path.normpath(os.path.realpath(path)).replace(os.sep, "/"): path
for path in [
*(self.base_models_roots or []),
*(self.embeddings_roots or []),
]
if isinstance(path, str) and path.strip() and os.path.exists(path)
}
for key, sub_type in OTHER_MODEL_FOLDER_SUBTYPES.items():
raw_paths = folder_path_map.get(key)
if not raw_paths:
continue
path_map = self._dedupe_existing_paths(raw_paths)
key_roots: List[str] = []
for real_path, business_path in sorted(
path_map.items(), key=lambda item: item[1].lower()
):
if real_path in seen_real_paths:
logger.warning(
"Detected the same folder '%s' under multiple other-model "
"categories ('%s' is already mapped). Keeping the first "
"category; please fix your path configuration.",
business_path,
seen_real_paths[real_path],
)
continue
seen_real_paths[real_path] = business_path
unique_paths.append(business_path)
key_roots.append(business_path)
sub_type_map[business_path] = sub_type
if real_path != business_path:
self.add_path_mapping(business_path, real_path)
covered_by = covered_real_paths.get(real_path)
if covered_by:
logger.warning(
"Detected an other-model root ('%s', category '%s') that "
"overlaps an existing checkpoints/embeddings root ('%s'). "
"The same files will appear on both pages; please review "
"your path configuration.",
business_path,
key,
covered_by,
)
if key_roots:
per_key_roots[key] = key_roots
return unique_paths, sub_type_map, per_key_roots
def _apply_library_paths( def _apply_library_paths(
self, self,
folder_paths: Mapping[str, Any], folder_paths: Mapping[str, Any],
@@ -1151,6 +1266,16 @@ class Config:
) = self._prepare_checkpoint_paths(checkpoint_paths, unet_paths) ) = self._prepare_checkpoint_paths(checkpoint_paths, unet_paths)
self.embeddings_roots = self._prepare_embedding_paths(embedding_paths) self.embeddings_roots = self._prepare_embedding_paths(embedding_paths)
other_path_map = {
key: folder_paths.get(key, []) or []
for key in self._get_enabled_other_folder_keys()
}
(
self.other_roots,
self.other_root_subtypes,
self.other_folder_roots,
) = self._prepare_other_paths(other_path_map)
# Process extra paths (only for LoRA Manager, not shared with ComfyUI) # Process extra paths (only for LoRA Manager, not shared with ComfyUI)
extra_paths = extra_folder_paths or {} extra_paths = extra_folder_paths or {}
extra_lora_paths = extra_paths.get("loras", []) or [] extra_lora_paths = extra_paths.get("loras", []) or []
@@ -1267,6 +1392,41 @@ class Config:
logger.warning(f"Error initializing embedding paths: {e}") logger.warning(f"Error initializing embedding paths: {e}")
return [] return []
def _init_other_paths(self) -> List[str]:
"""Initialize and validate other-model paths from ComfyUI settings.
Iterates the enabled OTHER_MODEL_FOLDER_SUBTYPES keys and pulls each
from ``folder_paths.get_folder_paths(key)`` (in standalone mode the
mock serves arbitrary keys from ``settings.json.folder_paths``).
"""
try:
folder_path_map: Dict[str, List[str]] = {}
for key in self._get_enabled_other_folder_keys():
try:
folder_path_map[key] = folder_paths.get_folder_paths(key)
except Exception as exc:
logger.debug("Error reading folder paths for '%s': %s", key, exc)
(
unique_paths,
self.other_root_subtypes,
self.other_folder_roots,
) = self._prepare_other_paths(folder_path_map)
logger.info(
"Found other model roots:"
+ ("\n - " + "\n - ".join(unique_paths) if unique_paths else "[]")
)
if not unique_paths:
logger.info("No valid other-model folders found in configuration")
return []
return unique_paths
except Exception as e:
logger.warning(f"Error initializing other model paths: {e}")
return []
def get_preview_static_url(self, preview_path: str) -> str: def get_preview_static_url(self, preview_path: str) -> str:
if not preview_path: if not preview_path:
return "" return ""
+7 -1
View File
@@ -219,6 +219,7 @@ class LoraManager:
lora_scanner = await ServiceRegistry.get_lora_scanner() lora_scanner = await ServiceRegistry.get_lora_scanner()
checkpoint_scanner = await ServiceRegistry.get_checkpoint_scanner() checkpoint_scanner = await ServiceRegistry.get_checkpoint_scanner()
embedding_scanner = await ServiceRegistry.get_embedding_scanner() embedding_scanner = await ServiceRegistry.get_embedding_scanner()
other_scanner = await ServiceRegistry.get_other_scanner()
# Initialize recipe scanner if needed # Initialize recipe scanner if needed
recipe_scanner = await ServiceRegistry.get_recipe_scanner() recipe_scanner = await ServiceRegistry.get_recipe_scanner()
@@ -236,6 +237,10 @@ class LoraManager:
embedding_scanner.initialize_in_background(), embedding_scanner.initialize_in_background(),
name="embedding_cache_init", name="embedding_cache_init",
), ),
asyncio.create_task(
other_scanner.initialize_in_background(),
name="other_cache_init",
),
asyncio.create_task( asyncio.create_task(
recipe_scanner.initialize_in_background(), name="recipe_cache_init" recipe_scanner.initialize_in_background(), name="recipe_cache_init"
), ),
@@ -328,6 +333,7 @@ class LoraManager:
all_roots.update(config.loras_roots) all_roots.update(config.loras_roots)
all_roots.update(config.base_models_roots or []) all_roots.update(config.base_models_roots or [])
all_roots.update(config.embeddings_roots or []) all_roots.update(config.embeddings_roots or [])
all_roots.update(config.other_roots or [])
total_deleted = 0 total_deleted = 0
total_size_freed = 0 total_size_freed = 0
@@ -460,7 +466,7 @@ class LoraManager:
# Cancel any in-flight scanner initialization tasks so thread-pool # Cancel any in-flight scanner initialization tasks so thread-pool
# workers (e.g. _initialize_cache_sync) can break out of their loops # workers (e.g. _initialize_cache_sync) can break out of their loops
# when the server shuts down (e.g. Ctrl+C on WSL). # when the server shuts down (e.g. Ctrl+C on WSL).
for name in ("lora_scanner", "checkpoint_scanner", "embedding_scanner"): for name in ("lora_scanner", "checkpoint_scanner", "embedding_scanner", "other_scanner"):
scanner = ServiceRegistry.get_service_sync(name) scanner = ServiceRegistry.get_service_sync(name)
if scanner is not None and hasattr(scanner, "cancel_task"): if scanner is not None and hasattr(scanner, "cancel_task"):
scanner.cancel_task() scanner.cancel_task()
+3 -2
View File
@@ -36,6 +36,7 @@ SCANNER_TYPE_MAP: dict[str, str] = {
"get_lora_scanner": "lora", "get_lora_scanner": "lora",
"get_checkpoint_scanner": "checkpoint", "get_checkpoint_scanner": "checkpoint",
"get_embedding_scanner": "embedding", "get_embedding_scanner": "embedding",
"get_other_scanner": "other",
} }
SCANNER_GETTER_NAMES = tuple(SCANNER_TYPE_MAP.keys()) SCANNER_GETTER_NAMES = tuple(SCANNER_TYPE_MAP.keys())
@@ -80,8 +81,8 @@ async def _find_scanner_for_model(
async def identify_model_type(model_path: str) -> str: async def identify_model_type(model_path: str) -> str:
"""Determine the model type (``\"lora\"``, ``\"checkpoint\"``, or """Determine the model type (``\"lora\"``, ``\"checkpoint\"``,
``\"embedding\"``) for *model_path*. ``\"embedding\"``, or ``\"other\"``) for *model_path*.
Falls back to ``\"lora\"`` when unknown. Falls back to ``\"lora\"`` when unknown.
""" """
+2
View File
@@ -658,6 +658,7 @@ class HealthCheckHandler:
"lora": ServiceRegistry.get_lora_scanner, "lora": ServiceRegistry.get_lora_scanner,
"checkpoint": ServiceRegistry.get_checkpoint_scanner, "checkpoint": ServiceRegistry.get_checkpoint_scanner,
"embedding": ServiceRegistry.get_embedding_scanner, "embedding": ServiceRegistry.get_embedding_scanner,
"other": ServiceRegistry.get_other_scanner,
"recipe": ServiceRegistry.get_recipe_scanner, "recipe": ServiceRegistry.get_recipe_scanner,
} }
@@ -757,6 +758,7 @@ class DoctorHandler:
("lora", "LoRAs", ServiceRegistry.get_lora_scanner), ("lora", "LoRAs", ServiceRegistry.get_lora_scanner),
("checkpoint", "Checkpoints", ServiceRegistry.get_checkpoint_scanner), ("checkpoint", "Checkpoints", ServiceRegistry.get_checkpoint_scanner),
("embedding", "Embeddings", ServiceRegistry.get_embedding_scanner), ("embedding", "Embeddings", ServiceRegistry.get_embedding_scanner),
("other", "Other Models", ServiceRegistry.get_other_scanner),
) )
) )
self._app_version_getter = app_version_getter self._app_version_getter = app_version_getter
@@ -35,6 +35,7 @@ _MODEL_TYPE_GETTER_NAMES: Dict[str, str] = {
"loras": "get_lora_scanner", "loras": "get_lora_scanner",
"checkpoints": "get_checkpoint_scanner", "checkpoints": "get_checkpoint_scanner",
"embeddings": "get_embedding_scanner", "embeddings": "get_embedding_scanner",
"other": "get_other_scanner",
} }
# Staged batch ids are ``uuid.uuid4().hex`` (32 lowercase hex chars). The id is # Staged batch ids are ``uuid.uuid4().hex`` (32 lowercase hex chars). The id is
+74
View File
@@ -0,0 +1,74 @@
import logging
from typing import Any, Dict
from aiohttp import web
from .base_model_routes import BaseModelRoutes
from .model_route_registrar import ModelRouteRegistrar
from ..services.other_model_service import OtherModelService
from ..services.service_registry import ServiceRegistry
from ..utils.constants import VALID_OTHER_CIVITAI_TYPES
logger = logging.getLogger(__name__)
class OtherRoutes(BaseModelRoutes):
"""Other-model-specific route controller (VAE, upscaler, text encoder, ...)"""
def __init__(self):
"""Initialize Other-model routes with OtherModel service"""
super().__init__()
self.template_name = "other.html"
async def initialize_services(self):
"""Initialize services from ServiceRegistry"""
other_scanner = await ServiceRegistry.get_other_scanner()
update_service = await ServiceRegistry.get_model_update_service()
self.service = OtherModelService(other_scanner, update_service=update_service)
self.set_model_update_service(update_service)
# Attach service dependencies
self.attach_service(self.service)
def setup_routes(self, app: web.Application, prefix: str = "other"):
"""Setup Other-model routes"""
# Schedule service initialization on app startup
app.on_startup.append(lambda _: self.initialize_services())
# Setup common routes with 'other' prefix (includes page route)
super().setup_routes(app, prefix)
def setup_specific_routes(self, registrar: ModelRouteRegistrar, prefix: str):
"""Setup Other-model-specific routes"""
# Other-model info by name
registrar.add_prefixed_route('GET', '/api/lm/{prefix}/info/{name}', prefix, self.get_other_model_info)
def _validate_civitai_model_type(self, model_type: str) -> bool:
"""Validate CivitAI model type for other models.
Accepts retired CivitAI types (CLIP, CLIPVision) as well grandfathered
models on CivitAI still carry them.
"""
return model_type.lower() in VALID_OTHER_CIVITAI_TYPES
def _get_expected_model_types(self) -> str:
"""Get expected model types string for error messages"""
return "VAE, Upscaler, TextEncoder, CLIPVision, Controlnet, or Other"
def _parse_specific_params(self, request: web.Request) -> Dict[str, Any]:
"""Parse other-model-specific parameters (none in Phase 1)."""
return {}
async def get_other_model_info(self, request: web.Request) -> web.Response:
"""Get detailed information for a specific other model by name"""
try:
name = request.match_info.get('name', '')
model_info = await self.service.get_model_info_by_name(name) # pyright: ignore[reportAttributeAccessIssue]
if model_info:
return web.json_response(model_info)
else:
return web.json_response({"error": "Model not found"}, status=404)
except Exception as e:
logger.error(f"Error in get_other_model_info: {e}", exc_info=True)
return web.json_response({"error": str(e)}, status=500)
+6 -1
View File
@@ -7,7 +7,7 @@ import logging
import os import os
import time import time
from ..utils.constants import VALID_LORA_SUB_TYPES, VALID_CHECKPOINT_SUB_TYPES from ..utils.constants import VALID_LORA_SUB_TYPES, VALID_CHECKPOINT_SUB_TYPES, VALID_OTHER_SUB_TYPES
from ..utils.models import BaseModelMetadata from ..utils.models import BaseModelMetadata
from ..utils.metadata_manager import MetadataManager from ..utils.metadata_manager import MetadataManager
from ..utils.usage_stats import UsageStats from ..utils.usage_stats import UsageStats
@@ -904,6 +904,11 @@ class BaseModelService(ABC):
and normalized_type not in VALID_CHECKPOINT_SUB_TYPES and normalized_type not in VALID_CHECKPOINT_SUB_TYPES
): ):
continue continue
if (
self.model_type == "other"
and normalized_type not in VALID_OTHER_SUB_TYPES
):
continue
type_counts[normalized_type] = type_counts.get(normalized_type, 0) + 1 type_counts[normalized_type] = type_counts.get(normalized_type, 0) + 1
+1
View File
@@ -68,6 +68,7 @@ PAGE_TYPE_MAP = {
'lora': 'loras', 'lora': 'loras',
'checkpoint': 'checkpoints', 'checkpoint': 'checkpoints',
'embedding': 'embeddings', 'embedding': 'embeddings',
'other': 'other',
} }
+6 -1
View File
@@ -118,13 +118,15 @@ class ModelServiceFactory:
def register_default_model_types(): def register_default_model_types():
"""Register the default model types (LoRA, Checkpoint, and Embedding)""" """Register the default model types (LoRA, Checkpoint, Embedding, and Other)"""
from ..services.lora_service import LoraService from ..services.lora_service import LoraService
from ..services.checkpoint_service import CheckpointService from ..services.checkpoint_service import CheckpointService
from ..services.embedding_service import EmbeddingService from ..services.embedding_service import EmbeddingService
from ..services.other_model_service import OtherModelService
from ..routes.lora_routes import LoraRoutes from ..routes.lora_routes import LoraRoutes
from ..routes.checkpoint_routes import CheckpointRoutes from ..routes.checkpoint_routes import CheckpointRoutes
from ..routes.embedding_routes import EmbeddingRoutes from ..routes.embedding_routes import EmbeddingRoutes
from ..routes.other_routes import OtherRoutes
# Register LoRA model type # Register LoRA model type
ModelServiceFactory.register_model_type('lora', LoraService, LoraRoutes) ModelServiceFactory.register_model_type('lora', LoraService, LoraRoutes)
@@ -134,3 +136,6 @@ def register_default_model_types():
# Register Embedding model type # Register Embedding model type
ModelServiceFactory.register_model_type('embedding', EmbeddingService, EmbeddingRoutes) ModelServiceFactory.register_model_type('embedding', EmbeddingService, EmbeddingRoutes)
# Register Other model type (VAE, upscaler, text encoder, ...)
ModelServiceFactory.register_model_type('other', OtherModelService, OtherRoutes)
+79
View File
@@ -0,0 +1,79 @@
import os
import logging
from typing import Any, Dict, Optional
from .base_model_service import BaseModelService
from .auto_tag_service import extract_auto_tags
from ..utils.models import OtherModelMetadata
from ..config import config
logger = logging.getLogger(__name__)
class OtherModelService(BaseModelService):
"""Other-model-specific service implementation (VAE, upscaler, text encoder, ...)"""
def __init__(self, scanner, update_service=None):
"""Initialize Other-model service
Args:
scanner: Other-model scanner instance
update_service: Optional service for remote update tracking.
"""
super().__init__("other", scanner, OtherModelMetadata, update_service=update_service)
async def format_response(self, model_data: Dict[str, Any]) -> Optional[Dict[str, Any]]:
"""Format other-model data for API response.
Returns None when the entry is missing critical fields (corrupted cache
row), so the handler layer can filter it out. See issue #730.
"""
# Guard against corrupted cache entries missing critical fields
file_path = model_data.get("file_path")
if not file_path or not isinstance(file_path, str):
logger.warning(
"Skipping corrupted other-model entry (missing file_path): %s",
model_data.get("file_name", "<unknown>"),
)
return None
# Get sub_type from cache entry (new canonical field)
sub_type = model_data.get("sub_type", "vae")
file_name = model_data.get("file_name") or ""
model_name = model_data.get("model_name") or file_name
folder = model_data.get("folder") or ""
return {
"model_name": model_name,
"file_name": file_name,
"preview_url": config.get_preview_static_url(model_data.get("preview_url", "")),
"preview_nsfw_level": model_data.get("preview_nsfw_level", 0),
"base_model": model_data.get("base_model", ""),
"folder": folder,
"sha256": model_data.get("sha256", ""),
"autov3": model_data.get("autov3"),
"file_path": file_path.replace(os.sep, "/"),
"file_size": model_data.get("size", 0),
"modified": model_data.get("modified", ""),
"tags": model_data.get("tags", []),
"from_civitai": model_data.get("from_civitai", True),
"notes": model_data.get("notes", ""),
"sub_type": sub_type,
"favorite": model_data.get("favorite", False),
"exclude": bool(model_data.get("exclude", False)),
"update_available": bool(model_data.get("update_available", False)),
"skip_metadata_refresh": bool(model_data.get("skip_metadata_refresh", False)),
"civitai": self.filter_civitai_data(model_data.get("civitai", {}), minimal=True),
"auto_tags": model_data.get("auto_tags") or extract_auto_tags(model_data),
"version_count": model_data.get("version_count"),
"hf_url": model_data.get("hf_url", ""),
}
def find_duplicate_hashes(self) -> Dict[str, Any]:
"""Find other models with duplicate SHA256 hashes"""
return self.scanner._hash_index.get_duplicate_hashes()
def find_duplicate_filenames(self) -> Dict[str, Any]:
"""Find other models with conflicting filenames"""
return self.scanner._hash_index.get_duplicate_filenames()
+468
View File
@@ -0,0 +1,468 @@
# pyright: reportImportCycles=false
# Lazy (function-local) imports still count as static edges in basedpyright's
# reportImportCycles, so the ServiceRegistry singleton pattern necessarily forms
# import cycles. Breaking them would require an architectural refactor.
import asyncio
import json
import logging
import os
from datetime import datetime
from typing import Any, Dict, List, Optional
from ..utils.models import OtherModelMetadata
from ..utils.file_utils import find_preview_file, normalize_path, calculate_autov3
from ..utils.metadata_manager import MetadataManager
from ..config import config
from .model_scanner import ModelScanner, _is_excluded_dir
from .model_hash_index import ModelHashIndex
logger = logging.getLogger(__name__)
class OtherScanner(ModelScanner):
"""Service for scanning and managing "other" model files.
Aggregates every enabled folder_paths category from
OTHER_MODEL_FOLDER_SUBTYPES (VAE, upscalers, text encoders, CLIP vision,
opt-in ControlNet) into one scanner; sub_type is derived from the root
containing the file (mirrors CheckpointScanner's checkpoints/unet split).
Hashing is lazy (checkpoint-style): text encoders can be ~10 GB, so the
initial scan records hash_status="pending" and the SHA256 is computed
on-demand via calculate_hash_for_model (e.g. when fetching CivitAI
metadata).
"""
def __init__(self):
# Same extension set as CheckpointScanner (ComfyUI's
# supported_pt_extensions plus ".gguf").
file_extensions = {
".ckpt",
".pt",
".pt2",
".bin",
".pth",
".safetensors",
".pkl",
".sft",
".gguf",
}
super().__init__(
model_type="other",
model_class=OtherModelMetadata,
file_extensions=file_extensions,
hash_index=ModelHashIndex(),
)
if not hasattr(self, "_hash_calculation_lock"):
self._hash_calculation_lock = asyncio.Lock()
self._hash_calculation_tasks: dict[str, asyncio.Task[Optional[str]]] = {}
async def _create_default_metadata(
self, file_path: str
) -> Optional[OtherModelMetadata]:
"""Create default metadata without calculating hash (lazy hash).
Other models include multi-GB text encoders, so hash calculation is
deferred until on-demand (e.g. CivitAI metadata fetch).
"""
try:
real_path = os.path.realpath(file_path)
if not os.path.exists(real_path):
logger.error(f"File not found: {file_path}")
return None
base_name = os.path.splitext(os.path.basename(file_path))[0]
dir_path = os.path.dirname(file_path)
# Find preview image
preview_url = find_preview_file(base_name, dir_path)
# AutoV3 reads only the safetensors header, so it is cheap even for
# large files; record the checked state at creation time ("" =
# checked but unavailable).
autov3 = calculate_autov3(real_path)
# Create metadata WITHOUT calculating hash
metadata = OtherModelMetadata(
file_name=base_name,
model_name=base_name,
file_path=normalize_path(file_path),
size=os.path.getsize(real_path),
modified=datetime.now().timestamp(),
sha256="", # Empty hash - will be calculated on-demand
base_model="Unknown",
preview_url=normalize_path(preview_url),
tags=[],
modelDescription="",
sub_type=self.resolve_sub_type_for_path(file_path) or "vae",
from_civitai=False, # Mark as local model since no hash yet
hash_status="pending", # Mark hash as pending
autov3=autov3 or "",
)
# Save the created metadata
logger.info(f"Creating other-model metadata (hash pending) for {file_path}")
await MetadataManager.save_metadata(file_path, metadata)
return metadata
except Exception as e:
logger.error(
f"Error creating default other-model metadata for {file_path}: {e}"
)
return None
async def calculate_hash_for_model(self, file_path: str) -> Optional[str]:
"""Calculate hash for a model on-demand with per-file singleflight.
Args:
file_path: Path to the model file
Returns:
SHA256 hash string, or None if calculation failed
"""
try:
real_path = os.path.realpath(file_path)
if not os.path.exists(real_path):
logger.error(f"File not found for hash calculation: {file_path}")
return None
metadata, _ = await MetadataManager.load_metadata(
file_path, self.model_class
)
if (
metadata is not None
and metadata.hash_status == "completed"
and metadata.sha256
):
# Ensure the in-memory hash index is populated even when
# the hash was already computed and persisted to the metadata
# file. Without this, usage tracking (and any other caller
# that queries get_hash_by_filename first) will miss on every
# lookup and keep calling back into this method, creating a
# tight loop that never populates the index.
self._hash_index.add_entry(
metadata.sha256.lower(),
file_path,
getattr(metadata, "autov3", None) or None,
)
return metadata.sha256
async with self._hash_calculation_lock:
metadata, _ = await MetadataManager.load_metadata(
file_path, self.model_class
)
if (
metadata is not None
and metadata.hash_status == "completed"
and metadata.sha256
):
self._hash_index.add_entry(
metadata.sha256.lower(),
file_path,
getattr(metadata, "autov3", None) or None,
)
return metadata.sha256
task = self._hash_calculation_tasks.get(real_path)
if task is None:
task = asyncio.create_task(
self._run_hash_calculation_task(file_path, real_path)
)
self._hash_calculation_tasks[real_path] = task
return await asyncio.shield(task)
except Exception as e:
logger.error(f"Error calculating hash for {file_path}: {e}")
return None
async def _run_hash_calculation_task(
self, file_path: str, real_path: str
) -> Optional[str]:
"""Run a hash calculation task and remove it from the in-flight map."""
try:
return await self._calculate_hash_for_model_uncached(file_path, real_path)
finally:
task = asyncio.current_task()
async with self._hash_calculation_lock:
if self._hash_calculation_tasks.get(real_path) is task:
del self._hash_calculation_tasks[real_path]
async def _calculate_hash_for_model_uncached(
self, file_path: str, real_path: str
) -> Optional[str]:
"""Calculate hash for a model without checking in-flight tasks."""
from ..utils.file_utils import calculate_sha256
try:
# Load current metadata
metadata, should_skip = await MetadataManager.load_metadata(
file_path, self.model_class
)
if metadata is None:
if should_skip:
logger.error(f"Invalid metadata found for {file_path}")
return None
created_metadata = await self._create_default_metadata(file_path)
if created_metadata is None:
logger.error(f"No metadata found for {file_path}")
return None
metadata = created_metadata
# Check if hash is already calculated
if metadata.hash_status == "completed" and metadata.sha256:
# Populate the in-memory hash index even for pre-computed
# hashes, mirroring the fix in calculate_hash_for_model.
self._hash_index.add_entry(
metadata.sha256.lower(),
file_path,
getattr(metadata, "autov3", None) or None,
)
return metadata.sha256
# Update status to calculating
metadata.hash_status = "calculating"
await MetadataManager.save_metadata(file_path, metadata)
# Calculate hash
logger.info(f"Calculating hash for other model: {file_path}")
sha256 = await calculate_sha256(real_path)
# Update metadata with hash
metadata.sha256 = sha256
metadata.hash_status = "completed"
await MetadataManager.save_metadata(file_path, metadata)
# Update hash index
self._hash_index.add_entry(
sha256.lower(),
file_path,
getattr(metadata, "autov3", None) or None,
)
# Update the in-memory cache entry so that subsequent
# _persist_current_cache / _save_persistent_cache calls
# write the hash back to the SQLite models table. Without
# this the hash only lives in the metadata file and the
# in-memory hash index, both of which are lost across
# restarts, causing the same re-computation loop on the
# next session.
if self._cache is not None and self._cache.raw_data:
for entry in self._cache.raw_data:
if entry.get("file_path") == file_path:
entry["sha256"] = sha256.lower()
entry["hash_status"] = "completed"
self.bump_cache_version()
break
logger.info(f"Hash calculated for other model: {file_path}")
return sha256
except Exception as e:
logger.error(f"Error calculating hash for {file_path}: {e}")
# Update status to failed
try:
metadata, _ = await MetadataManager.load_metadata(
file_path, self.model_class
)
if metadata:
metadata.hash_status = "failed"
await MetadataManager.save_metadata(file_path, metadata)
except Exception:
pass
return None
async def calculate_all_pending_hashes(
self, progress_callback=None
) -> Dict[str, int]:
"""Calculate hashes for all other models with pending hash status.
If cache is not initialized, scans filesystem directly for metadata files
with hash_status != 'completed'.
Args:
progress_callback: Optional callback(progress, total, current_file)
Returns:
Dict with 'completed', 'failed', 'total' counts
"""
# Try to get from cache first
cache = await self.get_cached_data()
if cache and cache.raw_data:
# Use cache if available
pending_models = [
item
for item in cache.raw_data
if item.get("hash_status") != "completed" or not item.get("sha256")
]
else:
# Cache not initialized, scan filesystem directly
pending_models = await self._find_pending_models_from_filesystem()
if not pending_models:
return {"completed": 0, "failed": 0, "total": 0}
total = len(pending_models)
completed = 0
failed = 0
for i, model_data in enumerate(pending_models):
file_path = model_data.get("file_path")
if not file_path:
continue
try:
sha256 = await self.calculate_hash_for_model(file_path)
if sha256:
completed += 1
else:
failed += 1
except Exception as e:
logger.error(f"Error calculating hash for {file_path}: {e}")
failed += 1
if progress_callback:
try:
await progress_callback(i + 1, total, file_path)
except Exception:
pass
return {"completed": completed, "failed": failed, "total": total}
async def _find_pending_models_from_filesystem(self) -> List[Dict[str, Any]]:
"""Scan filesystem for other-model metadata files with pending hash status."""
pending_models = []
for root_path in self.get_model_roots():
if not os.path.exists(root_path):
continue
for dirpath, dirnames, filenames in os.walk(root_path):
dirnames[:] = [d for d in dirnames if not _is_excluded_dir(d)]
for filename in filenames:
if not filename.endswith(".metadata.json"):
continue
metadata_path = os.path.join(dirpath, filename)
try:
with open(metadata_path, "r", encoding="utf-8") as f:
data = json.load(f)
# Check if hash is pending
hash_status = data.get("hash_status", "completed")
sha256 = data.get("sha256", "")
if hash_status != "completed" or not sha256:
# Find corresponding model file
model_name = filename.replace(".metadata.json", "")
model_path = None
# Look for model file with matching name
for ext in self.file_extensions:
potential_path = os.path.join(dirpath, model_name + ext)
if os.path.exists(potential_path):
model_path = potential_path
break
if model_path:
pending_models.append(
{
"file_path": model_path.replace(os.sep, "/"),
"hash_status": hash_status,
"sha256": sha256,
**{
k: v
for k, v in data.items()
if k
not in [
"file_path",
"hash_status",
"sha256",
]
},
}
)
except (json.JSONDecodeError, Exception) as e:
logger.debug(
f"Error reading metadata file {metadata_path}: {e}"
)
continue
return pending_models
def _root_sub_type_map(self) -> Dict[str, str]:
"""Return the configured business root -> sub_type map."""
root_map = getattr(config, "other_root_subtypes", None)
return root_map if isinstance(root_map, dict) else {}
def _resolve_sub_type(self, root_path: Optional[str]) -> Optional[str]:
"""Resolve the sub_type for a configured root path."""
if not root_path:
return None
normalized_root = self._normalize_path_value(root_path)
for root, sub_type in self._root_sub_type_map().items():
if self._normalize_path_value(root) == normalized_root:
return sub_type
return None
def resolve_sub_type_for_path(self, file_path: Optional[str]) -> Optional[str]:
"""Resolve sub_type from the configured root that contains the file.
Uses the longest-prefix match so nested roots (e.g. a controlnet root
inside a vae root) resolve to the most specific category.
"""
normalized_path = self._normalize_path_value(file_path)
if not normalized_path:
return None
best_length = 0
best_sub_type: Optional[str] = None
for root, sub_type in self._root_sub_type_map().items():
normalized_root = self._normalize_path_value(root)
if not normalized_root:
continue
if (
normalized_path == normalized_root
or normalized_path.startswith(f"{normalized_root}/")
) and len(normalized_root) > best_length:
best_length = len(normalized_root)
best_sub_type = sub_type
return best_sub_type
def adjust_metadata(self, metadata, file_path, root_path):
"""Adjust metadata during scanning to set sub_type."""
sub_type = self._resolve_sub_type(root_path) or self.resolve_sub_type_for_path(
file_path
)
if sub_type:
metadata.sub_type = sub_type
return metadata
def adjust_cached_entry(self, entry: Dict[str, Any]) -> Dict[str, Any]:
"""Adjust entries loaded from the persisted cache to ensure sub_type is set.
sub_type is location-derived: it is re-derived on cache load, never
trusted from the persisted snapshot.
"""
sub_type = self.resolve_sub_type_for_path(entry.get("file_path"))
if sub_type:
entry["sub_type"] = sub_type
return entry
def get_model_roots(self) -> List[str]:
"""Get other-model root directories"""
roots: List[str] = []
roots.extend(config.other_roots or [])
# Remove duplicates while preserving order
seen: set[str] = set()
unique_roots: List[str] = []
for root in roots:
if root and root not in seen:
seen.add(root)
unique_roots.append(root)
return unique_roots
+2
View File
@@ -59,6 +59,7 @@ _MODEL_TYPE_PAGE_MAP = {
"lora": "loras", "lora": "loras",
"checkpoint": "checkpoints", "checkpoint": "checkpoints",
"embedding": "embeddings", "embedding": "embeddings",
"other": "other",
} }
# Module-level alias so tests can spy on timer task creation without patching # Module-level alias so tests can spy on timer task creation without patching
@@ -983,6 +984,7 @@ class PendingDeleteService:
"get_lora_scanner", "get_lora_scanner",
"get_checkpoint_scanner", "get_checkpoint_scanner",
"get_embedding_scanner", "get_embedding_scanner",
"get_other_scanner",
): ):
getter = getattr(ServiceRegistry, getter_name, None) getter = getattr(ServiceRegistry, getter_name, None)
if not callable(getter): if not callable(getter):
+21
View File
@@ -314,6 +314,27 @@ class ServiceRegistry:
logger.debug(f"Created and registered {service_name}") logger.debug(f"Created and registered {service_name}")
return scanner return scanner
@classmethod
async def get_other_scanner(cls):
"""Get or create Other-model scanner instance"""
service_name = "other_scanner"
if service_name in cls._services:
return cls._services[service_name]
async with cls._get_lock(service_name):
# Double-check after acquiring lock
if service_name in cls._services:
return cls._services[service_name]
# Import here to avoid circular imports
from .other_scanner import OtherScanner
scanner = await OtherScanner.get_instance()
cls._services[service_name] = scanner
logger.debug(f"Created and registered {service_name}")
return scanner
@classmethod @classmethod
def clear_services(cls): def clear_services(cls):
"""Clear all registered services - mainly for testing""" """Clear all registered services - mainly for testing"""
+43
View File
@@ -83,6 +83,49 @@ VALID_LORA_SUB_TYPES = ["lora", "locon", "dora"]
VALID_CHECKPOINT_SUB_TYPES = ["checkpoint", "diffusion_model"] VALID_CHECKPOINT_SUB_TYPES = ["checkpoint", "diffusion_model"]
VALID_EMBEDDING_SUB_TYPES = ["embedding"] VALID_EMBEDDING_SUB_TYPES = ["embedding"]
# folder_paths key -> sub_type; single source of truth for extensibility.
# Adding support for a new ComfyUI folder category is a one-line change here.
OTHER_MODEL_FOLDER_SUBTYPES = {
"vae": "vae",
"upscale_models": "upscaler",
"text_encoders": "text_encoder",
"clip": "text_encoder", # legacy ComfyUI key
"clip_vision": "clip_vision",
"controlnet": "controlnet",
}
# folder_paths keys scanned by default; anything else in
# OTHER_MODEL_FOLDER_SUBTYPES (e.g. controlnet) is opt-in via the
# "enabled_other_folders" setting.
DEFAULT_OTHER_MODEL_FOLDERS = (
"vae",
"upscale_models",
"text_encoders",
"clip",
"clip_vision",
)
VALID_OTHER_SUB_TYPES = ["vae", "upscaler", "text_encoder", "clip_vision", "controlnet"]
# CivitAI model.type values accepted by the "other" page's fetch-metadata
# validation (lowercased). CLIP/CLIPVision are retired upstream but still
# appear on grandfathered models.
VALID_OTHER_CIVITAI_TYPES = {
"vae",
"upscaler",
"textencoder",
"clip",
"clipvision",
"controlnet",
"other",
}
# CivitAI model.type -> internal sub_type for the "other" model page.
CIVITAI_TYPE_TO_OTHER_SUB_TYPE = {
"vae": "vae",
"upscaler": "upscaler",
"textencoder": "text_encoder",
"clip": "text_encoder",
"clipvision": "clip_vision",
"controlnet": "controlnet",
}
# Backward compatibility alias # Backward compatibility alias
VALID_LORA_TYPES = VALID_LORA_SUB_TYPES VALID_LORA_TYPES = VALID_LORA_SUB_TYPES
+55 -1
View File
@@ -2,7 +2,7 @@ from dataclasses import dataclass, asdict, field
from typing import Callable, Dict, Optional, List, Any from typing import Callable, Dict, Optional, List, Any
from datetime import datetime from datetime import datetime
import os import os
from .constants import INVALID_AUTOV3_EMPTY_HASH from .constants import CIVITAI_TYPE_TO_OTHER_SUB_TYPE, INVALID_AUTOV3_EMPTY_HASH
from .model_utils import determine_base_model from .model_utils import determine_base_model
@@ -318,6 +318,60 @@ class CheckpointMetadata(BaseModelMetadata):
) )
@dataclass
class OtherModelMetadata(BaseModelMetadata):
"""Represents the metadata structure for an "other" model (VAE, upscaler,
text encoder, CLIP vision, ControlNet, ...).
The sub_type is location-derived: the OtherScanner sets it from the
folder_paths category whose root contains the file. The dataclass default
is only a placeholder.
"""
sub_type: str = "vae" # Placeholder; overridden by the scanner hooks
@classmethod
def from_civitai_info(
cls, version_info: Dict[str, Any], file_info: Dict[str, Any], save_path: str
) -> "OtherModelMetadata":
"""Create OtherModelMetadata instance from Civitai version info"""
file_name = file_info.get("name", "")
base_model = determine_base_model(version_info.get("baseModel", ""))
sha256_value = (file_info.get("hashes") or {}).get("SHA256", "").lower()
# Map the CivitAI model type onto our sub_types; unknown types keep the
# placeholder until the scanner re-derives sub_type from the location.
civitai_type = str(version_info.get("type", "") or "").lower()
sub_type = CIVITAI_TYPE_TO_OTHER_SUB_TYPE.get(civitai_type, "vae")
# Extract tags and description if available
tags = []
description = ""
model_data = version_info.get("model") or {}
if "tags" in model_data:
tags = model_data["tags"]
if "description" in model_data:
description = model_data["description"]
return cls(
file_name=os.path.splitext(file_name)[0],
model_name=model_data.get("name", os.path.splitext(file_name)[0]),
file_path=save_path.replace(os.sep, "/"),
size=file_info.get("sizeKB", 0) * 1024,
modified=datetime.now().timestamp(),
sha256=sha256_value,
base_model=base_model,
preview_url="", # Will be updated after preview download
preview_nsfw_level=0,
from_civitai=True,
civitai=version_info,
sub_type=sub_type,
tags=tags,
modelDescription=description,
# Direct read: the downloaded file IS file_info, no SHA256 matching.
autov3=normalize_autov3((file_info.get("hashes") or {}).get("AutoV3")),
)
@dataclass @dataclass
class EmbeddingMetadata(BaseModelMetadata): class EmbeddingMetadata(BaseModelMetadata):
"""Represents the metadata structure for an Embedding model""" """Represents the metadata structure for an Embedding model"""
+12
View File
@@ -17,6 +17,18 @@
"embeddings": [ "embeddings": [
"C:/path/to/your/embeddings_folder", "C:/path/to/your/embeddings_folder",
"C:/path/to/another/embeddings_folder" "C:/path/to/another/embeddings_folder"
],
"vae": [
"C:/path/to/your/vae_folder"
],
"upscale_models": [
"C:/path/to/your/upscale_models_folder"
],
"text_encoders": [
"C:/path/to/your/text_encoders_folder"
],
"clip_vision": [
"C:/path/to/your/clip_vision_folder"
] ]
}, },
"auto_organize_exclusions": [] "auto_organize_exclusions": []
+277
View File
@@ -0,0 +1,277 @@
"""Tests for other-model path handling in py/config.py."""
import logging
import os
import pytest
from py import config as config_module
from py.services.settings_manager import get_settings_manager
def _normalize(path: str) -> str:
return os.path.normpath(path).replace(os.sep, "/")
def _make_config(**overrides) -> config_module.Config:
"""Create a bare Config instance for _prepare_other_paths tests."""
config = config_module.Config.__new__(config_module.Config)
config._path_mappings = {}
config._preview_root_paths = set()
config._cached_fingerprint = None
config.base_models_roots = []
config.embeddings_roots = []
for key, value in overrides.items():
setattr(config, key, value)
return config
class TestPrepareOtherPaths:
"""Unit tests for Config._prepare_other_paths."""
def test_maps_each_folder_key_to_sub_type(self, tmp_path):
roots = {
"vae": tmp_path / "vae",
"upscale_models": tmp_path / "upscale_models",
"text_encoders": tmp_path / "text_encoders",
"clip": tmp_path / "clip",
"clip_vision": tmp_path / "clip_vision",
"controlnet": tmp_path / "controlnet",
}
for root in roots.values():
root.mkdir()
config = _make_config()
unique, sub_type_map, per_key = config._prepare_other_paths(
{key: [str(root)] for key, root in roots.items()}
)
assert len(unique) == 6
assert sub_type_map[_normalize(str(roots["vae"]))] == "vae"
assert sub_type_map[_normalize(str(roots["upscale_models"]))] == "upscaler"
assert sub_type_map[_normalize(str(roots["text_encoders"]))] == "text_encoder"
# Legacy ComfyUI 'clip' key maps to text_encoder as well
assert sub_type_map[_normalize(str(roots["clip"]))] == "text_encoder"
assert sub_type_map[_normalize(str(roots["clip_vision"]))] == "clip_vision"
assert sub_type_map[_normalize(str(roots["controlnet"]))] == "controlnet"
assert per_key["vae"] == [_normalize(str(roots["vae"]))]
assert per_key["controlnet"] == [_normalize(str(roots["controlnet"]))]
def test_missing_or_unknown_keys_are_skipped(self, tmp_path):
vae_root = tmp_path / "vae"
vae_root.mkdir()
config = _make_config()
unique, sub_type_map, per_key = config._prepare_other_paths(
{
"vae": [str(vae_root)],
"does_not_exist_key": [str(tmp_path / "nope_dir")],
"upscale_models": [],
}
)
assert unique == [_normalize(str(vae_root))]
assert set(per_key.keys()) == {"vae"}
def test_nonexistent_paths_are_filtered(self, tmp_path):
config = _make_config()
unique, _, _ = config._prepare_other_paths(
{"vae": [str(tmp_path / "missing_vae")]}
)
assert unique == []
def test_cross_category_overlap_warns_and_keeps_first(
self, tmp_path, caplog
):
"""The same physical folder under two categories warns; first wins."""
shared = tmp_path / "shared"
shared.mkdir()
config = _make_config()
with caplog.at_level(logging.WARNING, logger=config_module.logger.name):
unique, sub_type_map, per_key = config._prepare_other_paths(
{
"vae": [str(shared)],
"upscale_models": [str(shared)],
}
)
assert unique == [_normalize(str(shared))]
assert sub_type_map[_normalize(str(shared))] == "vae"
assert "upscale_models" not in per_key
warnings = [
record.message
for record in caplog.records
if record.levelname == "WARNING"
and "multiple other-model categories" in record.message
]
assert len(warnings) == 1
def test_cross_scanner_overlap_warns_but_keeps_path(self, tmp_path, caplog):
"""An other root overlapping a checkpoint root warns but stays managed."""
shared = tmp_path / "shared_models"
shared.mkdir()
config = _make_config(base_models_roots=[_normalize(str(shared))])
with caplog.at_level(logging.WARNING, logger=config_module.logger.name):
unique, sub_type_map, _ = config._prepare_other_paths(
{"vae": [str(shared)]}
)
# Kept on purpose: dropping would silently unmanage the files
assert unique == [_normalize(str(shared))]
assert sub_type_map[_normalize(str(shared))] == "vae"
warnings = [
record.message
for record in caplog.records
if record.levelname == "WARNING"
and "overlaps an existing checkpoints/embeddings root" in record.message
]
assert len(warnings) == 1
def test_no_warning_for_disjoint_roots(self, tmp_path, caplog):
checkpoints_root = tmp_path / "checkpoints"
checkpoints_root.mkdir()
vae_root = tmp_path / "vae"
vae_root.mkdir()
config = _make_config(base_models_roots=[_normalize(str(checkpoints_root))])
with caplog.at_level(logging.WARNING, logger=config_module.logger.name):
unique, _, _ = config._prepare_other_paths({"vae": [str(vae_root)]})
assert unique == [_normalize(str(vae_root))]
warnings = [
record.message
for record in caplog.records
if record.levelname == "WARNING" and "overlap" in record.message.lower()
]
assert warnings == []
class TestInitOtherPaths:
"""Config._init_other_paths with mocked folder_paths (plugin + standalone modes).
Config only depends on ``folder_paths.get_folder_paths(key)``: ComfyUI in
plugin mode, or MockFolderPaths serving ``settings.json.folder_paths`` in
standalone mode. A dict-backed stub therefore covers both.
"""
def _stub_folder_paths(self, monkeypatch, mapping):
def get_folder_paths(key):
value = mapping.get(key, [])
return [value] if isinstance(value, str) else list(value)
monkeypatch.setattr(
config_module.folder_paths, "get_folder_paths", get_folder_paths
)
def test_default_enabled_keys_exclude_controlnet(self, monkeypatch, tmp_path):
dirs = {}
for key in (
"vae",
"upscale_models",
"text_encoders",
"clip",
"clip_vision",
"controlnet",
):
path = tmp_path / key
path.mkdir()
dirs[key] = str(path)
self._stub_folder_paths(monkeypatch, dirs)
config = _make_config()
roots = config._init_other_paths()
assert _normalize(dirs["controlnet"]) not in roots
assert _normalize(dirs["controlnet"]) not in config.other_root_subtypes
for key in ("vae", "upscale_models", "text_encoders", "clip", "clip_vision"):
assert _normalize(dirs[key]) in roots
def test_controlnet_opt_in_via_setting(self, monkeypatch, tmp_path):
controlnet_dir = tmp_path / "controlnet"
controlnet_dir.mkdir()
self._stub_folder_paths(monkeypatch, {"controlnet": str(controlnet_dir)})
get_settings_manager().set("enabled_other_folders", ["controlnet"])
config = _make_config()
roots = config._init_other_paths()
assert _normalize(str(controlnet_dir)) in roots
assert (
config.other_root_subtypes[_normalize(str(controlnet_dir))]
== "controlnet"
)
def test_unknown_opt_in_keys_are_ignored(self, monkeypatch, tmp_path):
vae_dir = tmp_path / "vae"
vae_dir.mkdir()
self._stub_folder_paths(monkeypatch, {"vae": str(vae_dir)})
get_settings_manager().set("enabled_other_folders", ["not_a_real_key", 42])
config = _make_config()
roots = config._init_other_paths()
assert roots == [_normalize(str(vae_dir))]
def test_apply_library_paths_picks_up_other_keys(self, monkeypatch, tmp_path):
vae_dir = tmp_path / "vae"
vae_dir.mkdir()
config = _make_config()
monkeypatch.setattr(config, "_initialize_symlink_mappings", lambda: None)
config._apply_library_paths(
{
"loras": [],
"checkpoints": [],
"unet": [],
"embeddings": [],
"vae": [str(vae_dir)],
}
)
assert config.other_roots == [_normalize(str(vae_dir))]
assert config.other_root_subtypes == {
_normalize(str(vae_dir)): "vae"
}
assert config.other_folder_roots == {"vae": [_normalize(str(vae_dir))]}
class TestOtherRootsWiring:
"""other_roots participates in symlink and preview root bookkeeping."""
def test_symlink_roots_include_other_roots(self):
config = _make_config()
config.loras_roots = ["/loras"]
config.embeddings_roots = ["/embeddings"]
config.other_roots = ["/vae"]
config.extra_loras_roots = []
config.extra_checkpoints_roots = []
config.extra_unet_roots = []
config.extra_embeddings_roots = []
assert "/vae" in config._symlink_roots()
def test_preview_roots_include_other_roots(self, tmp_path):
vae_dir = tmp_path / "vae"
vae_dir.mkdir()
config = _make_config()
config.loras_roots = []
config.embeddings_roots = []
config.other_roots = [_normalize(str(vae_dir))]
config.extra_loras_roots = []
config.extra_checkpoints_roots = []
config.extra_unet_roots = []
config.extra_embeddings_roots = []
config.recipes_path = ""
config._rebuild_preview_roots()
assert config.is_preview_path_allowed(str(vae_dir / "model.preview.png"))
+4 -2
View File
@@ -132,6 +132,7 @@ async def test_lora_manager_lifecycle(monkeypatch: pytest.MonkeyPatch, tmp_path:
"lora": _DummyScanner("lora"), "lora": _DummyScanner("lora"),
"checkpoint": _DummyScanner("checkpoint"), "checkpoint": _DummyScanner("checkpoint"),
"embedding": _DummyScanner("embedding"), "embedding": _DummyScanner("embedding"),
"other": _DummyScanner("other"),
"recipe": _DummyScanner("recipe"), "recipe": _DummyScanner("recipe"),
} }
@@ -147,6 +148,7 @@ async def test_lora_manager_lifecycle(monkeypatch: pytest.MonkeyPatch, tmp_path:
monkeypatch.setattr(lora_manager.ServiceRegistry, "get_lora_scanner", lambda: _stub("lora_scanner", scanners["lora"])) monkeypatch.setattr(lora_manager.ServiceRegistry, "get_lora_scanner", lambda: _stub("lora_scanner", scanners["lora"]))
monkeypatch.setattr(lora_manager.ServiceRegistry, "get_checkpoint_scanner", lambda: _stub("checkpoint_scanner", scanners["checkpoint"])) monkeypatch.setattr(lora_manager.ServiceRegistry, "get_checkpoint_scanner", lambda: _stub("checkpoint_scanner", scanners["checkpoint"]))
monkeypatch.setattr(lora_manager.ServiceRegistry, "get_embedding_scanner", lambda: _stub("embedding_scanner", scanners["embedding"])) monkeypatch.setattr(lora_manager.ServiceRegistry, "get_embedding_scanner", lambda: _stub("embedding_scanner", scanners["embedding"]))
monkeypatch.setattr(lora_manager.ServiceRegistry, "get_other_scanner", lambda: _stub("other_scanner", scanners["other"]))
monkeypatch.setattr(lora_manager.ServiceRegistry, "get_recipe_scanner", lambda: _stub("recipe_scanner", scanners["recipe"])) monkeypatch.setattr(lora_manager.ServiceRegistry, "get_recipe_scanner", lambda: _stub("recipe_scanner", scanners["recipe"]))
migration_calls: list[bool] = [] migration_calls: list[bool] = []
@@ -205,7 +207,7 @@ async def test_lora_manager_lifecycle(monkeypatch: pytest.MonkeyPatch, tmp_path:
await asyncio.gather(*pending) await asyncio.gather(*pending)
task_names = {task.get_name() for task in scheduled_tasks} task_names = {task.get_name() for task in scheduled_tasks}
assert {"lora_cache_init", "checkpoint_cache_init", "embedding_cache_init", "recipe_cache_init", "post_init_tasks", "cleanup_bak_files"}.issubset(task_names) assert {"lora_cache_init", "checkpoint_cache_init", "embedding_cache_init", "other_cache_init", "recipe_cache_init", "post_init_tasks", "cleanup_bak_files"}.issubset(task_names)
# Startup sweep: an expired pending-delete purge task is spawned during # Startup sweep: an expired pending-delete purge task is spawned during
# service initialization (covers both plugin and standalone modes). # service initialization (covers both plugin and standalone modes).
@@ -219,4 +221,4 @@ async def test_lora_manager_lifecycle(monkeypatch: pytest.MonkeyPatch, tmp_path:
for root in (loras_root, checkpoints_root, embeddings_root): for root in (loras_root, checkpoints_root, embeddings_root):
assert not any(path.suffix == ".bak" for path in root.rglob("*")), f"Backup files remain in {root}" assert not any(path.suffix == ".bak" for path in root.rglob("*")), f"Backup files remain in {root}"
assert {"civitai_client", "download_manager", "websocket_manager", "lora_scanner", "checkpoint_scanner", "embedding_scanner", "recipe_scanner"}.issubset(registry_calls) assert {"civitai_client", "download_manager", "websocket_manager", "lora_scanner", "checkpoint_scanner", "embedding_scanner", "other_scanner", "recipe_scanner"}.issubset(registry_calls)
+118
View File
@@ -0,0 +1,118 @@
import json
import pytest
from aiohttp import web
from py.routes.other_routes import OtherRoutes
from py.services.other_model_service import OtherModelService
class DummyRequest:
def __init__(self, *, match_info=None):
self.match_info = match_info or {}
class StubOtherModelService:
def __init__(self):
self.info = {}
async def get_model_info_by_name(self, name):
value = self.info.get(name)
if isinstance(value, Exception):
raise value
return value
@pytest.fixture
def routes():
handler = OtherRoutes()
handler.service = StubOtherModelService() # pyright: ignore[reportAttributeAccessIssue]
return handler
def test_common_and_specific_routes_registered():
"""Registration smoke test: /api/lm/other/* surface plus the /other page."""
app = web.Application()
OtherRoutes().setup_routes(app)
registered = {(route.method, route.resource.canonical) for route in app.router.routes()}
assert ("GET", "/other") in registered
assert ("GET", "/api/lm/other/list") in registered
assert ("GET", "/api/lm/other/model-types") in registered
assert ("GET", "/api/lm/other/roots") in registered
assert ("POST", "/api/lm/other/fetch-civitai") in registered
assert ("POST", "/api/lm/other/delete") in registered
assert ("POST", "/api/lm/other/move_model") in registered
assert ("GET", "/api/lm/other/info/{name}") in registered
def test_template_name_is_other_page():
assert OtherRoutes().template_name == "other.html"
@pytest.mark.parametrize(
"model_type",
["VAE", "Upscaler", "TextEncoder", "CLIP", "CLIPVision", "Controlnet", "Other"],
)
def test_validate_civitai_model_type_accepts_other_types(model_type):
assert OtherRoutes()._validate_civitai_model_type(model_type) is True
@pytest.mark.parametrize("model_type", ["Lora", "Checkpoint", "TextualInversion"])
def test_validate_civitai_model_type_rejects_foreign_types(model_type):
assert OtherRoutes()._validate_civitai_model_type(model_type) is False
def test_get_expected_model_types_mentions_supported_types():
expected = OtherRoutes()._get_expected_model_types()
for name in ("VAE", "Upscaler", "TextEncoder", "CLIPVision", "Controlnet"):
assert name in expected
async def test_get_other_model_info_success(routes):
routes.service.info["demo"] = {"name": "demo"}
response = await routes.get_other_model_info(DummyRequest(match_info={"name": "demo"}))
payload = json.loads(response.text)
assert payload == {"name": "demo"}
async def test_get_other_model_info_missing(routes):
response = await routes.get_other_model_info(DummyRequest(match_info={"name": "missing"}))
payload = json.loads(response.text)
assert response.status == 404
assert payload == {"error": "Model not found"}
async def test_get_other_model_info_error(routes):
routes.service.info["demo"] = RuntimeError("boom")
response = await routes.get_other_model_info(DummyRequest(match_info={"name": "demo"}))
payload = json.loads(response.text)
assert response.status == 500
assert payload == {"error": "boom"}
@pytest.mark.asyncio
async def test_initialize_services_builds_other_model_service(monkeypatch):
from py.services.service_registry import ServiceRegistry
sentinel_scanner = object()
sentinel_update_service = object()
async def fake_scanner():
return sentinel_scanner
async def fake_update_service():
return sentinel_update_service
monkeypatch.setattr(ServiceRegistry, "get_other_scanner", staticmethod(fake_scanner))
monkeypatch.setattr(
ServiceRegistry, "get_model_update_service", staticmethod(fake_update_service)
)
handler = OtherRoutes()
await handler.initialize_services()
assert isinstance(handler.service, OtherModelService)
assert handler.service.model_type == "other"
assert handler.service.scanner is sentinel_scanner
+337
View File
@@ -0,0 +1,337 @@
"""Tests for OtherScanner: root aggregation, sub_type derivation, lazy hash."""
import asyncio
import json
import os
from pathlib import Path
from unittest.mock import MagicMock, patch
import pytest
from py import config as config_module
from py.services import model_scanner
from py.services.model_scanner import ModelScanner
from py.services.other_scanner import OtherScanner
from py.utils.models import OtherModelMetadata
def _normalize(path) -> str:
return str(path).replace(os.sep, "/")
@pytest.fixture(autouse=True)
def reset_model_scanner_singletons():
ModelScanner._instances.clear()
ModelScanner._locks.clear()
yield
ModelScanner._instances.clear()
ModelScanner._locks.clear()
@pytest.fixture
def other_config(monkeypatch, tmp_path):
"""Point the global config at a synthetic set of other-model roots."""
vae_root = tmp_path / "vae"
upscaler_root = tmp_path / "upscale_models"
te_root = tmp_path / "text_encoders"
clip_root = tmp_path / "clip"
clip_vision_root = tmp_path / "clip_vision"
for root in (vae_root, upscaler_root, te_root, clip_root, clip_vision_root):
root.mkdir()
roots = [
_normalize(vae_root),
_normalize(upscaler_root),
_normalize(te_root),
_normalize(clip_root),
_normalize(clip_vision_root),
]
subtypes = {
_normalize(vae_root): "vae",
_normalize(upscaler_root): "upscaler",
_normalize(te_root): "text_encoder",
_normalize(clip_root): "text_encoder",
_normalize(clip_vision_root): "clip_vision",
}
monkeypatch.setattr(config_module.config, "other_roots", roots)
monkeypatch.setattr(config_module.config, "other_root_subtypes", subtypes)
return {
"roots": roots,
"subtypes": subtypes,
"vae": _normalize(vae_root),
"upscaler": _normalize(upscaler_root),
"text_encoders": _normalize(te_root),
"clip": _normalize(clip_root),
"clip_vision": _normalize(clip_vision_root),
}
def _make_scanner() -> OtherScanner:
"""Create a scanner without __init__ to avoid async initialization."""
scanner = object.__new__(OtherScanner)
scanner.model_type = "other"
scanner.model_class = OtherModelMetadata
scanner.file_extensions = {".safetensors", ".pt", ".bin"}
scanner._hash_index = MagicMock()
return scanner
class TestOtherScannerRoots:
"""Root aggregation and sub_type resolution."""
def test_get_model_roots_aggregates_and_dedupes(self, other_config, monkeypatch):
scanner = _make_scanner()
monkeypatch.setattr(
config_module.config,
"other_roots",
other_config["roots"] + [other_config["vae"]],
)
roots = scanner.get_model_roots()
assert roots == other_config["roots"]
def test_get_model_roots_empty_when_unconfigured(self, monkeypatch):
scanner = _make_scanner()
monkeypatch.setattr(config_module.config, "other_roots", None)
assert scanner.get_model_roots() == []
def test_resolve_sub_type_for_each_default_category(self, other_config):
scanner = _make_scanner()
cases = [
(other_config["vae"], "vae"),
(other_config["upscaler"], "upscaler"),
(other_config["text_encoders"], "text_encoder"),
# Legacy ComfyUI 'clip' key maps to text_encoder as well
(other_config["clip"], "text_encoder"),
(other_config["clip_vision"], "clip_vision"),
]
for root, expected in cases:
file_path = f"{root}/model.safetensors"
assert scanner.resolve_sub_type_for_path(file_path) == expected
def test_resolve_sub_type_longest_prefix_wins(self, monkeypatch, tmp_path):
"""A nested root (controlnet inside vae) resolves to the inner category."""
outer = tmp_path / "vae"
inner = outer / "controlnet"
inner.mkdir(parents=True)
monkeypatch.setattr(
config_module.config,
"other_root_subtypes",
{_normalize(outer): "vae", _normalize(inner): "controlnet"},
)
scanner = _make_scanner()
assert (
scanner.resolve_sub_type_for_path(f"{_normalize(inner)}/cn.safetensors")
== "controlnet"
)
assert (
scanner.resolve_sub_type_for_path(f"{_normalize(outer)}/vae.safetensors")
== "vae"
)
def test_resolve_sub_type_none_for_unknown_or_empty(self, other_config):
scanner = _make_scanner()
assert scanner.resolve_sub_type_for_path(None) is None
assert scanner.resolve_sub_type_for_path("") is None
assert scanner.resolve_sub_type_for_path("/unrelated/model.safetensors") is None
def test_adjust_metadata_sets_sub_type(self, other_config):
scanner = _make_scanner()
metadata = OtherModelMetadata(
file_name="te",
model_name="te",
file_path=f"{other_config['text_encoders']}/te.safetensors",
size=1,
modified=0.0,
sha256="",
base_model="Unknown",
preview_url="",
)
result = scanner.adjust_metadata(
metadata,
metadata.file_path,
other_config["text_encoders"],
)
assert result.sub_type == "text_encoder"
def test_adjust_cached_entry_rederives_sub_type(self, other_config):
"""Persisted sub_type is never trusted: it is re-derived from location."""
scanner = _make_scanner()
entry = {
"file_path": f"{other_config['clip']}/legacy.safetensors",
"sub_type": "vae", # stale value from an old snapshot
}
result = scanner.adjust_cached_entry(entry)
assert result["sub_type"] == "text_encoder"
def test_adjust_cached_entry_keeps_value_when_root_unknown(self, other_config):
scanner = _make_scanner()
entry = {
"file_path": "/gone/model.safetensors",
"sub_type": "upscaler",
}
result = scanner.adjust_cached_entry(entry)
assert result["sub_type"] == "upscaler"
class TestOtherScannerLazyHash:
"""Lazy hashing: pending by default, singleflight on-demand calculation."""
@pytest.mark.asyncio
async def test_default_metadata_has_pending_hash(self, other_config):
vae_file = Path(other_config["vae"]) / "vae_model.safetensors"
vae_file.write_text("fake vae content", encoding="utf-8")
scanner = OtherScanner()
metadata = await scanner._create_default_metadata(_normalize(vae_file))
assert metadata is not None
assert metadata.sha256 == ""
assert metadata.hash_status == "pending"
assert metadata.from_civitai is False
assert metadata.sub_type == "vae"
@pytest.mark.asyncio
async def test_default_metadata_sub_type_from_location(self, other_config):
te_file = Path(other_config["text_encoders"]) / "t5.safetensors"
te_file.write_text("fake text encoder", encoding="utf-8")
scanner = OtherScanner()
metadata = await scanner._create_default_metadata(_normalize(te_file))
assert metadata is not None
assert metadata.sub_type == "text_encoder"
@pytest.mark.asyncio
async def test_calculate_hash_for_model_completes_pending(self, other_config):
model_file = Path(other_config["upscaler"]) / "upscaler.safetensors"
model_file.write_text("fake upscaler content", encoding="utf-8")
normalized_file = _normalize(model_file)
scanner = OtherScanner()
metadata = await scanner._create_default_metadata(normalized_file)
assert metadata is not None and metadata.hash_status == "pending"
hash_result = await scanner.calculate_hash_for_model(normalized_file)
assert hash_result is not None
assert len(hash_result) == 64
metadata_file = model_file.with_suffix(".metadata.json")
saved_data = json.loads(metadata_file.read_text(encoding="utf-8"))
assert saved_data["sha256"] == hash_result
assert saved_data["hash_status"] == "completed"
@pytest.mark.asyncio
async def test_calculate_hash_singleflight_same_file(self, other_config):
"""Concurrent calls for the same file share one SHA256 task."""
model_file = Path(other_config["vae"]) / "shared.safetensors"
model_file.write_text("fake content", encoding="utf-8")
normalized_file = _normalize(model_file)
real_file = os.path.realpath(normalized_file)
scanner = OtherScanner()
metadata = await scanner._create_default_metadata(normalized_file)
assert metadata is not None
calls = []
async def fake_calculate_sha256(file_path: str) -> str:
calls.append(file_path)
await asyncio.sleep(0.01)
return "a" * 64
with patch(
"py.utils.file_utils.calculate_sha256", side_effect=fake_calculate_sha256
):
results = await asyncio.gather(
*[scanner.calculate_hash_for_model(normalized_file) for _ in range(8)]
)
assert calls == [real_file]
assert results == ["a" * 64] * 8
assert scanner._hash_calculation_tasks == {}
@pytest.mark.asyncio
async def test_calculate_hash_skips_completed(self, other_config):
model_file = Path(other_config["clip_vision"]) / "cv.safetensors"
model_file.write_text("fake content", encoding="utf-8")
normalized_file = _normalize(model_file)
scanner = OtherScanner()
metadata = await scanner._create_default_metadata(normalized_file)
assert metadata is not None
# Simulate an already-completed hash
metadata.sha256 = "existing_hash"
metadata.hash_status = "completed"
from py.utils.metadata_manager import MetadataManager
await MetadataManager.save_metadata(normalized_file, metadata)
with patch("py.utils.file_utils.calculate_sha256") as mock_calc:
hash_result = await scanner.calculate_hash_for_model(normalized_file)
assert hash_result == "existing_hash"
mock_calc.assert_not_called()
@pytest.mark.asyncio
async def test_calculate_all_pending_hashes(self, other_config):
for index in range(3):
model_file = Path(other_config["vae"]) / f"model_{index}.safetensors"
model_file.write_text(f"content {index}", encoding="utf-8")
scanner = OtherScanner()
for index in range(3):
model_file = Path(other_config["vae"]) / f"model_{index}.safetensors"
await scanner._create_default_metadata(_normalize(model_file))
progress_calls = []
async def progress_callback(current, total, file_path):
progress_calls.append((current, total, file_path))
result = await scanner.calculate_all_pending_hashes(progress_callback)
assert result["total"] == 3
assert result["completed"] == 3
assert result["failed"] == 0
assert len(progress_calls) == 3
class TestOtherModelMetadataFromCivitai:
"""CivitAI type mapping in OtherModelMetadata.from_civitai_info."""
def _build(self, civitai_type: str) -> OtherModelMetadata:
return OtherModelMetadata.from_civitai_info(
{
"type": civitai_type,
"baseModel": "SDXL",
"model": {"name": "Model", "tags": ["tag"], "description": "desc"},
},
{"name": "model.safetensors", "sizeKB": 1, "hashes": {"SHA256": "AB"}},
"/tmp/model.safetensors",
)
@pytest.mark.parametrize(
"civitai_type,expected",
[
("VAE", "vae"),
("Upscaler", "upscaler"),
("TextEncoder", "text_encoder"),
("CLIP", "text_encoder"),
("CLIPVision", "clip_vision"),
("Controlnet", "controlnet"),
("Other", "vae"), # unknown types fall back to the placeholder
],
)
def test_civitai_type_mapping(self, civitai_type, expected):
metadata = self._build(civitai_type)
assert metadata.sub_type == expected
assert metadata.sha256 == "ab"
assert metadata.tags == ["tag"]
def test_page_type_maps_to_other():
"""The WS progress page type for the other scanner is 'other'."""
assert model_scanner.PAGE_TYPE_MAP["other"] == "other"
scanner = _make_scanner()
assert scanner.page_type == "other"
@@ -7,6 +7,7 @@ from unittest.mock import MagicMock, AsyncMock
from py.services.lora_service import LoraService from py.services.lora_service import LoraService
from py.services.checkpoint_service import CheckpointService from py.services.checkpoint_service import CheckpointService
from py.services.embedding_service import EmbeddingService from py.services.embedding_service import EmbeddingService
from py.services.other_model_service import OtherModelService
class TestLoraServiceFormatResponse: class TestLoraServiceFormatResponse:
@@ -206,6 +207,89 @@ class TestEmbeddingServiceFormatResponse:
assert "model_type" not in result # Removed in refactoring assert "model_type" not in result # Removed in refactoring
class TestOtherModelServiceFormatResponse:
"""Test OtherModelService.format_response includes sub_type."""
@pytest.fixture
def mock_scanner(self):
scanner = MagicMock()
scanner._hash_index = MagicMock()
return scanner
@pytest.fixture
def other_service(self, mock_scanner):
return OtherModelService(mock_scanner)
@pytest.mark.asyncio
async def test_format_response_includes_sub_type(self, other_service):
"""format_response should include sub_type field."""
other_data = {
"model_name": "Test VAE",
"file_name": "test_vae",
"preview_url": "test.webp",
"preview_nsfw_level": 0,
"base_model": "SDXL",
"folder": "",
"sha256": "abc123",
"file_path": "/models/vae/test_vae.safetensors",
"size": 1000,
"modified": 1234567890.0,
"tags": [],
"from_civitai": True,
"notes": "",
"favorite": False,
"sub_type": "vae",
"civitai": {},
}
result = await other_service.format_response(other_data)
assert "sub_type" in result
assert result["sub_type"] == "vae"
assert "model_type" not in result # Removed in refactoring
@pytest.mark.asyncio
async def test_format_response_defaults_to_vae(self, other_service):
"""format_response should default to 'vae' if no sub_type field."""
other_data = {
"model_name": "Test Upscaler",
"file_name": "test_upscaler",
"preview_url": "test.webp",
"preview_nsfw_level": 0,
"base_model": "SD1.5",
"folder": "",
"sha256": "abc123",
"file_path": "/models/upscale_models/test.pth",
"size": 1000,
"modified": 1234567890.0,
"tags": [],
"from_civitai": True,
"civitai": {},
}
result = await other_service.format_response(other_data)
assert result["sub_type"] == "vae"
assert "model_type" not in result # Removed in refactoring
@pytest.mark.asyncio
async def test_format_response_returns_none_on_missing_file_path(self, other_service):
"""format_response returns None when file_path is missing (corrupted row)."""
other_data = {
"model_name": "Test",
"file_name": "test",
"file_path": None, # corrupted: missing file_path
"folder": "",
"sha256": "abc",
"tags": [],
"from_civitai": True,
"civitai": {},
"sub_type": "text_encoder",
}
result = await other_service.format_response(other_data)
assert result is None
class TestFormatResponseCorruptedEntries: class TestFormatResponseCorruptedEntries:
"""Test format_response handles corrupted cache entries gracefully (issue #730). """Test format_response handles corrupted cache entries gracefully (issue #730).