diff --git a/py/recipes/base.py b/py/recipes/base.py index 14f2cc62..3e2364b4 100644 --- a/py/recipes/base.py +++ b/py/recipes/base.py @@ -11,7 +11,7 @@ import re from typing import Dict, List, Any, Optional, Tuple from abc import ABC, abstractmethod from ..config import config -from ..utils.constants import VALID_LORA_TYPES, VALID_CHECKPOINT_SUB_TYPES +from ..utils.constants import MODEL_WEIGHT_FILE_TYPES, VALID_LORA_TYPES, VALID_CHECKPOINT_SUB_TYPES from ..utils.civitai_utils import rewrite_preview_url logger = logging.getLogger(__name__) @@ -155,9 +155,9 @@ class RecipeMetadataParser(ABC): # Process file information if available if 'files' in civitai_info: - # Find the primary model file (type="Model" and primary=true) in the files list + # Find the primary model file (weights-type and primary=true) in the files list model_file = next((file for file in civitai_info.get('files', []) - if file.get('type') == 'Model' and file.get('primary') == True), None) + if file.get('type') in MODEL_WEIGHT_FILE_TYPES and file.get('primary') == True), None) if model_file: # Get size @@ -261,11 +261,21 @@ class RecipeMetadataParser(ABC): checkpoint['id'] = civitai_data.get('id', 0) if 'files' in civitai_data: + # Prefer the file CivitAI marked primary; fall back to any + # weights-type file (providers without primary flags). model_file = next( ( file for file in civitai_data.get('files', []) - if file.get('type') == 'Model' + if file.get('type') in MODEL_WEIGHT_FILE_TYPES + and file.get('primary') is True + ), + None, + ) or next( + ( + file + for file in civitai_data.get('files', []) + if file.get('type') in MODEL_WEIGHT_FILE_TYPES ), None, ) diff --git a/py/routes/base_model_routes.py b/py/routes/base_model_routes.py index 2ebed727..4d2e842d 100644 --- a/py/routes/base_model_routes.py +++ b/py/routes/base_model_routes.py @@ -30,6 +30,7 @@ from ..services.websocket_progress_callback import ( WebSocketProgressCallback, ) from ..utils.exif_utils import ExifUtils +from ..utils.constants import MODEL_WEIGHT_FILE_TYPES from ..utils.metadata_manager import MetadataManager from .model_route_registrar import COMMON_ROUTE_DEFINITIONS, ModelRouteRegistrar from .handlers.model_handlers import ( @@ -251,7 +252,7 @@ class BaseModelRoutes(ABC): def _find_model_file(self, files): """Find the appropriate model file from the files list - can be overridden by subclasses.""" - return next((file for file in files if file.get("type") in ("Model", "Diffusion Model") and file.get("primary") is True), None) + return next((file for file in files if file.get("type") in MODEL_WEIGHT_FILE_TYPES and file.get("primary") is True), None) def get_handler(self, name: str) -> Callable[[web.Request], Awaitable[web.StreamResponse]]: """Expose handlers for subclasses or tests.""" diff --git a/py/services/civitai_client.py b/py/services/civitai_client.py index d084ff07..81f2f297 100644 --- a/py/services/civitai_client.py +++ b/py/services/civitai_client.py @@ -21,6 +21,7 @@ from .model_metadata_provider import ( from .downloader import get_downloader from .errors import RateLimitError, ResourceNotFoundError from ..utils.civitai_utils import resolve_license_payload +from ..utils.constants import MODEL_WEIGHT_FILE_TYPES logger = logging.getLogger(__name__) @@ -538,10 +539,16 @@ class CivitaiClient: return model_versions[0] def _extract_primary_model_hash(self, version_entry: Dict[str, Any]) -> Optional[str]: + # Prefer the generic "Model" file (most reliable version identity); + # fall back to any other weights-type primary. for file_info in version_entry.get("files", []): if file_info.get("type") == "Model" and file_info.get("primary"): - hashes = file_info.get("hashes", {}) - model_hash = hashes.get("SHA256") + model_hash = (file_info.get("hashes", {}) or {}).get("SHA256") + if model_hash: + return model_hash + for file_info in version_entry.get("files", []): + if file_info.get("type") in MODEL_WEIGHT_FILE_TYPES and file_info.get("primary"): + model_hash = (file_info.get("hashes", {}) or {}).get("SHA256") if model_hash: return model_hash return None diff --git a/py/services/download_manager.py b/py/services/download_manager.py index 67efc5fe..29d8cf5d 100644 --- a/py/services/download_manager.py +++ b/py/services/download_manager.py @@ -18,6 +18,7 @@ from ..utils.models import LoraMetadata, CheckpointMetadata, EmbeddingMetadata from ..utils.constants import ( CARD_PREVIEW_WIDTH, DIFFUSION_MODEL_BASE_MODELS, + MODEL_WEIGHT_FILE_TYPES, SUPPORTED_DOWNLOAD_SKIP_BASE_MODELS, VALID_LORA_TYPES, ) @@ -46,6 +47,11 @@ CIVITAI_DOWNLOAD_URL_PREFIXES = ( ) +# File types that are never the intended download target even when CivitAI +# marks them primary — configs/archives/workflows are auxiliary artifacts. +NON_DOWNLOADABLE_PRIMARY_TYPES = ("Config", "Archive", "Workflow", "Training Data") + + class DownloadManager: _instance = None _lock = asyncio.Lock() @@ -1500,7 +1506,7 @@ class DownloadManager: f for f in files if f.get("primary") - and f.get("type") in ("Model", "Negative", "Diffusion Model", "UNet") + and f.get("type") in MODEL_WEIGHT_FILE_TYPES ), None, ) @@ -1540,21 +1546,52 @@ class DownloadManager: # Fallback to primary file if no match found if not file_info: logger.debug("[download] Looking for primary file as fallback") + # Prefer a weights-type file CivitAI marked primary; then any + # weights-type file (providers without primary flags, e.g. + # civarchive); then trust CivitAI's primary flag regardless of + # type — newer types like 'Enhancement LoRA' are valid primary + # files. Weights files are preferred over non-weights primary + # files so a Config/Archive primary never replaces a Model. file_info = next( ( f for f in files - if f.get("primary") and f.get("type") in ("Model", "Negative", "Diffusion Model", "UNet") + if f.get("primary") and f.get("type") in MODEL_WEIGHT_FILE_TYPES ), None, ) if file_info: logger.debug( - "[download] Fallback primary file selected: id=%s, name=%s", + "[download] Fallback primary file selected (primary + weights): id=%s, name=%s", file_info.get("id"), file_info.get("name"), ) else: - logger.debug("[download] No primary file found in fallback lookup") + file_info = next( + (f for f in files if f.get("type") in MODEL_WEIGHT_FILE_TYPES), + None, + ) + if file_info: + logger.debug( + "[download] Fallback primary file selected (weights type, no primary flag): id=%s, name=%s", + file_info.get("id"), file_info.get("name"), + ) + else: + file_info = next( + ( + f + for f in files + if f.get("primary") + and f.get("type") not in NON_DOWNLOADABLE_PRIMARY_TYPES + ), + None, + ) + if file_info: + logger.debug( + "[download] Fallback primary file selected (trusting CivitAI primary flag): id=%s, name=%s, type=%s", + file_info.get("id"), file_info.get("name"), file_info.get("type"), + ) + else: + logger.debug("[download] No primary file found in fallback lookup") if not file_info: return {"success": False, "error": "No suitable file found in metadata"} diff --git a/py/utils/constants.py b/py/utils/constants.py index 134cf875..974537a2 100644 --- a/py/utils/constants.py +++ b/py/utils/constants.py @@ -62,6 +62,20 @@ MODEL_FILE_EXTENSIONS = { ".gguf", } +# CivitAI ModelFile.type values eligible as the main download file. +# Mirrors CivitAI's getPrimaryFile() (model-helpers.ts): weight types are +# preferred, but any file CivitAI marks `primary` is accepted — newer types +# like 'Enhancement LoRA' (Anima/AIR image-editing LoRAs) are valid primary +# files despite not being in the traditional weights allowlist. +MODEL_WEIGHT_FILE_TYPES = ( + "Model", + "Pruned Model", + "Negative", + "UNet", + "Diffusion Model", + "Enhancement LoRA", +) + # Valid sub-types for each scanner type VALID_LORA_SUB_TYPES = ["lora", "locon", "dora"] VALID_CHECKPOINT_SUB_TYPES = ["checkpoint", "diffusion_model"] diff --git a/static/js/components/ContextMenu/RecipeContextMenu.js b/static/js/components/ContextMenu/RecipeContextMenu.js index ac181572..49be1425 100644 --- a/static/js/components/ContextMenu/RecipeContextMenu.js +++ b/static/js/components/ContextMenu/RecipeContextMenu.js @@ -1,6 +1,7 @@ import { BaseContextMenu } from './BaseContextMenu.js'; import { ModelContextMenuMixin } from './ModelContextMenuMixin.js'; import { showToast, copyToClipboard, sendLoraToWorkflow } from '../../utils/uiHelpers.js'; +import { isModelWeightFile } from '../../utils/modelFileTypes.js'; import { setSessionItem, removeSessionItem } from '../../utils/storageHelpers.js'; import { updateRecipeMetadata } from '../../api/recipeApi.js'; import { state } from '../../state/index.js'; @@ -255,7 +256,7 @@ export class RecipeContextMenu extends BaseContextMenu { loras: validLoras.map(lora => { const civitaiInfo = lora.civitaiInfo; const modelFile = civitaiInfo.files ? - civitaiInfo.files.find(file => file.type === 'Model') : null; + civitaiInfo.files.find(file => isModelWeightFile(file.type)) : null; return { // Basic lora info diff --git a/static/js/components/RecipeModal.js b/static/js/components/RecipeModal.js index ec3e409b..18e31544 100644 --- a/static/js/components/RecipeModal.js +++ b/static/js/components/RecipeModal.js @@ -1,5 +1,6 @@ // Recipe Modal Component import { showToast, copyToClipboard, sendLoraToWorkflow, sendModelPathToWorkflow, openCivitaiByMetadata, stripLoraTags, sendPromptToWorkflow, sendGenParamsToWorkflow } from '../utils/uiHelpers.js'; +import { isModelWeightFile } from '../utils/modelFileTypes.js'; import { translate } from '../utils/i18nHelpers.js'; import { state } from '../state/index.js'; import { setSessionItem, removeSessionItem, getStorageItem, setStorageItem } from '../utils/storageHelpers.js'; @@ -1412,7 +1413,7 @@ class RecipeModal { loras: validLoras.map(lora => { const civitaiInfo = lora.civitaiInfo; const modelFile = civitaiInfo.files ? - civitaiInfo.files.find(file => file.type === 'Model') : null; + civitaiInfo.files.find(file => isModelWeightFile(file.type)) : null; return { // Basic lora info diff --git a/static/js/managers/DownloadManager.js b/static/js/managers/DownloadManager.js index 37879dc3..2cb5b50b 100644 --- a/static/js/managers/DownloadManager.js +++ b/static/js/managers/DownloadManager.js @@ -3,6 +3,7 @@ import { showToast, setupAutoNewlineOnPaste } from '../utils/uiHelpers.js'; import { state } from '../state/index.js'; import { LoadingManager } from './LoadingManager.js'; import { getModelApiClient, resetAndReload } from '../api/modelApiFactory.js'; +import { isModelWeightFile } from '../utils/modelFileTypes.js'; import { getStorageItem, setStorageItem } from '../utils/storageHelpers.js'; import { FolderTreeManager } from '../components/FolderTreeManager.js'; import { translate } from '../utils/i18nHelpers.js'; @@ -557,8 +558,7 @@ export class DownloadManager { const firstImage = version.images?.find(img => !img.url.endsWith('.mp4')); const thumbnailUrl = firstImage ? firstImage.url : '/loras_static/images/no-preview.png'; - // Count model-type files per version - const modelFiles = (version.files || []).filter(f => f.type === 'Model' || f.type === 'UNet' || f.type === 'Diffusion Model'); + const modelFiles = (version.files || []).filter(f => isModelWeightFile(f.type)); const primaryFile = modelFiles.find(f => f.primary) || modelFiles[0] || {}; const fileSize = version.modelSizeKB ? (version.modelSizeKB / 1024).toFixed(2) : @@ -685,7 +685,7 @@ export class DownloadManager { if (!version) return; this.currentVersion = version; - const modelFiles = (version.files || []).filter(f => f.type === 'Model' || f.type === 'UNet' || f.type === 'Diffusion Model'); + const modelFiles = (version.files || []).filter(f => isModelWeightFile(f.type)); document.getElementById('versionStep').style.display = 'none'; document.getElementById('fileSelectionStep').style.display = 'block'; @@ -747,7 +747,7 @@ export class DownloadManager { return; } - const modelFiles = (version.files || []).filter(f => f.type === 'Model' || f.type === 'UNet' || f.type === 'Diffusion Model'); + const modelFiles = (version.files || []).filter(f => isModelWeightFile(f.type)); this.selectedFile = modelFiles.find(f => f.id.toString() === selectedRadio.value); console.log('[download] confirmFileSelection: selected file id=%s, name="%s", type="%s", metadata=%o', diff --git a/static/js/utils/modelFileTypes.js b/static/js/utils/modelFileTypes.js new file mode 100644 index 00000000..76b53068 --- /dev/null +++ b/static/js/utils/modelFileTypes.js @@ -0,0 +1,15 @@ +// CivitAI ModelFile.type values eligible as the main download file. +// Mirrors the backend constant MODEL_WEIGHT_FILE_TYPES (py/utils/constants.py). +// Keep both lists in sync when CivitAI introduces new file types. +export const MODEL_WEIGHT_FILE_TYPES = [ + 'Model', + 'Pruned Model', + 'Negative', + 'UNet', + 'Diffusion Model', + 'Enhancement LoRA', +]; + +export function isModelWeightFile(type) { + return MODEL_WEIGHT_FILE_TYPES.includes(type); +} diff --git a/tests/services/test_download_manager_basic.py b/tests/services/test_download_manager_basic.py index 2c3a8a01..cc98794a 100644 --- a/tests/services/test_download_manager_basic.py +++ b/tests/services/test_download_manager_basic.py @@ -123,10 +123,7 @@ def metadata_provider(monkeypatch): class DummyProvider: def __init__(self): self.calls = [] - - async def get_model_version(self, model_id, model_version_id): - self.calls.append((model_id, model_version_id)) - return { + self.payload = { "id": 42, "model": {"type": "LoRA", "tags": ["fantasy"]}, "baseModel": "BaseModel", @@ -141,6 +138,10 @@ def metadata_provider(monkeypatch): ], } + async def get_model_version(self, model_id, model_version_id): + self.calls.append((model_id, model_version_id)) + return self.payload + provider = DummyProvider() monkeypatch.setattr( download_manager, @@ -233,6 +234,165 @@ async def test_successful_download_uses_defaults( assert captured["download_urls"] == ["https://example.invalid/file.safetensors"] +@pytest.mark.asyncio +async def test_download_accepts_enhancement_lora_primary_file( + monkeypatch, scanners, metadata_provider, tmp_path +): + """A version whose only file has type 'Enhancement LoRA' (Anima/AIR + image-editing LoRAs) must download — previously failed with + "No suitable file found in metadata" because the type was missing from + the primary-file weights allowlist.""" + manager = DownloadManager() + metadata_provider.payload = { + "id": 3219121, + "model": {"type": "LORA", "tags": ["style"]}, + "baseModel": "Anima", + "creator": {"username": "Deskup"}, + "files": [ + { + "id": 3100968, + "type": "Enhancement LoRA", + "primary": True, + "name": "deskup-anima-edit-general.safetensors", + "sizeKB": 358501.13, + "downloadUrl": "https://example.invalid/deskup-anima-edit-general.safetensors", + } + ], + } + + captured = {} + + async def fake_execute_download(self, **kwargs): + captured.update( + { + "download_urls": kwargs["download_urls"], + "model_type": kwargs["model_type"], + } + ) + return {"success": True} + + monkeypatch.setattr( + DownloadManager, "_execute_download", fake_execute_download, raising=False + ) + + result = await manager.download_from_civitai( + model_id=2850692, + model_version_id=3219121, + save_dir=str(tmp_path), + use_default_paths=True, + progress_callback=None, + source=None, + ) + + assert result["success"] is True + assert captured["model_type"] == "lora" + assert captured["download_urls"] == [ + "https://example.invalid/deskup-anima-edit-general.safetensors" + ] + + +@pytest.mark.asyncio +async def test_download_falls_back_to_civitai_primary_flag_regardless_of_type( + monkeypatch, scanners, metadata_provider, tmp_path +): + """If no weights-type file exists, trust CivitAI's `primary` flag on any + file — mirrors CivitAI's getPrimaryFile() which never excludes a file by + type.""" + manager = DownloadManager() + metadata_provider.payload = { + "id": 77, + "model": {"type": "LORA", "tags": ["concept"]}, + "baseModel": "Anima", + "creator": {"username": "Author"}, + "files": [ + { + "id": 100, + "type": "Other", + "primary": True, + "name": "custom-type-lora.safetensors", + "downloadUrl": "https://example.invalid/custom-type-lora.safetensors", + } + ], + } + + captured = {} + + async def fake_execute_download(self, **kwargs): + captured["download_urls"] = kwargs["download_urls"] + return {"success": True} + + monkeypatch.setattr( + DownloadManager, "_execute_download", fake_execute_download, raising=False + ) + + result = await manager.download_from_civitai( + model_version_id=77, + save_dir=str(tmp_path), + use_default_paths=True, + progress_callback=None, + source=None, + ) + + assert result["success"] is True + assert captured["download_urls"] == [ + "https://example.invalid/custom-type-lora.safetensors" + ] + + +@pytest.mark.asyncio +async def test_download_prefers_weights_file_over_non_weights_primary( + monkeypatch, scanners, metadata_provider, tmp_path +): + """A Config/Archive-type primary must never replace an existing weights + file — the weights file wins even without the primary flag.""" + manager = DownloadManager() + metadata_provider.payload = { + "id": 78, + "model": {"type": "LORA", "tags": ["concept"]}, + "baseModel": "BaseModel", + "creator": {"username": "Author"}, + "files": [ + { + "id": 201, + "type": "Config", + "primary": True, + "name": "config.json", + "downloadUrl": "https://example.invalid/config.json", + }, + { + "id": 202, + "type": "Model", + "primary": False, + "name": "weights.safetensors", + "downloadUrl": "https://example.invalid/weights.safetensors", + }, + ], + } + + captured = {} + + async def fake_execute_download(self, **kwargs): + captured["download_urls"] = kwargs["download_urls"] + return {"success": True} + + monkeypatch.setattr( + DownloadManager, "_execute_download", fake_execute_download, raising=False + ) + + result = await manager.download_from_civitai( + model_version_id=78, + save_dir=str(tmp_path), + use_default_paths=True, + progress_callback=None, + source=None, + ) + + assert result["success"] is True + assert captured["download_urls"] == [ + "https://example.invalid/weights.safetensors" + ] + + @pytest.mark.asyncio async def test_download_keeps_save_dir_when_use_save_dir_as_root( monkeypatch, scanners, metadata_provider, tmp_path