diff --git a/py/services/download_coordinator.py b/py/services/download_coordinator.py index dbbe2e17..2b0f6466 100644 --- a/py/services/download_coordinator.py +++ b/py/services/download_coordinator.py @@ -83,6 +83,7 @@ class DownloadCoordinator: save_dir=payload.get("model_root"), relative_path=payload.get("relative_path", ""), use_default_paths=payload.get("use_default_paths", False), + use_save_dir_as_root=payload.get("use_save_dir_as_root", False), progress_callback=progress_callback, download_id=download_id, source=payload.get("source"), diff --git a/py/services/download_manager.py b/py/services/download_manager.py index e9b40a9e..67efc5fe 100644 --- a/py/services/download_manager.py +++ b/py/services/download_manager.py @@ -217,6 +217,7 @@ class DownloadManager: download_id: str | None = None, source: str | None = None, file_params: Dict[str, Any] | None = None, + use_save_dir_as_root: bool = False, ) -> Dict[str, Any]: """Download model from Civitai with task tracking and concurrency control @@ -257,6 +258,7 @@ class DownloadManager: "save_dir": save_dir, "relative_path": relative_path, "use_default_paths": bool(use_default_paths), + "use_save_dir_as_root": bool(use_save_dir_as_root), "source": source, "file_params": copy.deepcopy(file_params) if file_params is not None else None, "progress": 0, @@ -287,6 +289,7 @@ class DownloadManager: use_default_paths, source, file_params, + use_save_dir_as_root, ) ) @@ -321,6 +324,7 @@ class DownloadManager: use_default_paths: bool = False, source: str | None = None, file_params: Dict[str, Any] | None = None, + use_save_dir_as_root: bool = False, ): """Execute download with semaphore to limit concurrency""" # Update status to waiting @@ -401,6 +405,7 @@ class DownloadManager: ), source, file_params, + use_save_dir_as_root=use_save_dir_as_root, ) # Update status based on result @@ -621,6 +626,7 @@ class DownloadManager: "save_dir": info.get("save_dir"), "relative_path": info.get("relative_path", ""), "use_default_paths": bool(info.get("use_default_paths", False)), + "use_save_dir_as_root": bool(info.get("use_save_dir_as_root", False)), "source": info.get("source"), "file_params": copy.deepcopy(info.get("file_params")), "transfer_backend": info.get("transfer_backend", "aria2"), @@ -643,6 +649,7 @@ class DownloadManager: "save_dir": record.get("save_dir"), "relative_path": record.get("relative_path", ""), "use_default_paths": bool(record.get("use_default_paths", False)), + "use_save_dir_as_root": bool(record.get("use_save_dir_as_root", False)), "source": record.get("source"), "file_params": copy.deepcopy(record.get("file_params")), "progress": record.get("progress", 0), @@ -1001,6 +1008,7 @@ class DownloadManager: bool(restored.get("use_default_paths", False)), restored.get("source"), restored.get("file_params"), + bool(restored.get("use_save_dir_as_root", False)), ) ) continue @@ -1134,6 +1142,7 @@ class DownloadManager: transfer_backend: str = "python", source: str | None = None, file_params: Dict[str, Any] | None = None, + use_save_dir_as_root: bool = False, ) -> Dict[str, Any]: """Wrapper for original download_from_civitai implementation""" try: @@ -1362,36 +1371,41 @@ class DownloadManager: # Handle use_default_paths if use_default_paths: settings_manager = get_settings_manager() - # Set save_dir based on model type - if model_type == "checkpoint": - if is_diffusion_model: - default_path = settings_manager.get("default_unet_root") - error_msg = "Default unet root path not set in settings" - else: - default_path = settings_manager.get("default_checkpoint_root") - error_msg = "Default checkpoint root path not set in settings" - if not default_path: - return { - "success": False, - "error": error_msg, - } - save_dir = default_path - elif model_type == "lora": - default_path = settings_manager.get("default_lora_root") - if not default_path: - return { - "success": False, - "error": "Default lora root path not set in settings", - } - save_dir = default_path - elif model_type == "embedding": - default_path = settings_manager.get("default_embedding_root") - if not default_path: - return { - "success": False, - "error": "Default embedding root path not set in settings", - } - save_dir = default_path + # With use_save_dir_as_root, an explicitly provided save_dir is kept + # as the base root and the path template is resolved underneath it. + # Otherwise fall back to the configured default root, which keeps the + # classic "download to default root" behavior for regular downloads. + if not save_dir or not use_save_dir_as_root: + # Set save_dir based on model type + if model_type == "checkpoint": + if is_diffusion_model: + default_path = settings_manager.get("default_unet_root") + error_msg = "Default unet root path not set in settings" + else: + default_path = settings_manager.get("default_checkpoint_root") + error_msg = "Default checkpoint root path not set in settings" + if not default_path: + return { + "success": False, + "error": error_msg, + } + save_dir = default_path + elif model_type == "lora": + default_path = settings_manager.get("default_lora_root") + if not default_path: + return { + "success": False, + "error": "Default lora root path not set in settings", + } + save_dir = default_path + elif model_type == "embedding": + default_path = settings_manager.get("default_embedding_root") + if not default_path: + return { + "success": False, + "error": "Default embedding root path not set in settings", + } + save_dir = default_path # Calculate relative path using template relative_path = self._calculate_relative_path(version_info, model_type) @@ -2761,6 +2775,7 @@ class DownloadManager: bool(persisted.get("use_default_paths", False)), persisted.get("source"), persisted.get("file_params"), + bool(persisted.get("use_save_dir_as_root", False)), ), ) except Exception as exc: diff --git a/static/js/api/baseModelApi.js b/static/js/api/baseModelApi.js index e276d70c..6fd07250 100644 --- a/static/js/api/baseModelApi.js +++ b/static/js/api/baseModelApi.js @@ -1233,7 +1233,7 @@ export class BaseModelApiClient { } } - async downloadModel(modelId, versionId, modelRoot, relativePath, useDefaultPaths = false, downloadId, source = null, fileParams = null) { + async downloadModel(modelId, versionId, modelRoot, relativePath, useDefaultPaths = false, downloadId, source = null, fileParams = null, useSaveDirAsRoot = false) { try { const response = await fetch(DOWNLOAD_ENDPOINTS.download, { method: 'POST', @@ -1244,6 +1244,7 @@ export class BaseModelApiClient { model_root: modelRoot, relative_path: relativePath, use_default_paths: useDefaultPaths, + use_save_dir_as_root: useSaveDirAsRoot, download_id: downloadId, ...(source ? { source } : {}), ...(fileParams ? { file_params: fileParams } : {}) diff --git a/static/js/components/shared/ModelVersionsTab.js b/static/js/components/shared/ModelVersionsTab.js index 98ef3c1a..3bb30890 100644 --- a/static/js/components/shared/ModelVersionsTab.js +++ b/static/js/components/shared/ModelVersionsTab.js @@ -1307,15 +1307,41 @@ export function initVersionsTab({ }); } - async function resolveDownloadPathFromCurrentVersion() { + function getCurrentInLibraryVersion() { if (!normalizedCurrentVersionId || !controller.record?.versions) { return null; } - - const currentVersion = controller.record.versions.find( + return controller.record.versions.find( v => v.versionId === normalizedCurrentVersionId && v.isInLibrary && v.filePath - ); - if (!currentVersion?.filePath) { + ) || null; + } + + function getDownloadPathTemplate() { + try { + const singularType = modelType.replace(/s$/, ''); + const templates = state.global?.settings?.download_path_templates; + return (templates && templates[singularType]) || ''; + } catch (error) { + return ''; + } + } + + function shouldResolveTemplatePath(targetVersion, pathInfo) { + if (!getDownloadPathTemplate() || !pathInfo?.modelRoot) { + return false; + } + const currentVersion = getCurrentInLibraryVersion(); + const currentBase = normalizeBaseModelName(currentVersion?.baseModel); + const targetBase = normalizeBaseModelName(targetVersion?.baseModel); + if (!currentBase || !targetBase || currentBase === targetBase) { + return false; + } + return true; + } + + async function resolveDownloadPathFromCurrentVersion() { + const currentVersion = getCurrentInLibraryVersion(); + if (!currentVersion) { return null; } @@ -1372,10 +1398,13 @@ export function initVersionsTab({ try { const pathInfo = await resolveDownloadPathFromCurrentVersion(); + const resolveTemplatePath = shouldResolveTemplatePath(version, pathInfo); const success = await downloadManager.downloadVersionWithDefaults(modelType, modelId, versionId, { versionName: version.name || `#${version.versionId}`, modelRoot: pathInfo?.modelRoot || '', - targetFolder: pathInfo?.targetFolder || '', + targetFolder: resolveTemplatePath ? '' : (pathInfo?.targetFolder || ''), + useDefaultPaths: resolveTemplatePath ? true : null, + useSaveDirAsRoot: resolveTemplatePath, }); if (success) { diff --git a/static/js/managers/DownloadManager.js b/static/js/managers/DownloadManager.js index 71578d26..37879dc3 100644 --- a/static/js/managers/DownloadManager.js +++ b/static/js/managers/DownloadManager.js @@ -912,6 +912,7 @@ export class DownloadManager { modelRoot = '', targetFolder = '', useDefaultPaths = false, + useSaveDirAsRoot = false, source = null, fileParams = null, closeModal = false, @@ -923,7 +924,7 @@ export class DownloadManager { } const displayName = versionName || `#${versionId}`; - const retryParams = { modelId, versionId, versionName, modelRoot, targetFolder, useDefaultPaths, source, fileParams, closeModal: false }; + const retryParams = { modelId, versionId, versionName, modelRoot, targetFolder, useDefaultPaths, useSaveDirAsRoot, source, fileParams, closeModal: false }; let ws = null; let updateProgress = () => { }; let cancelled = false; @@ -995,7 +996,8 @@ export class DownloadManager { useDefaultPaths, downloadId, source, - fileParams + fileParams, + useSaveDirAsRoot ); if (cancelled) { @@ -1809,7 +1811,9 @@ export class DownloadManager { versionName = '', source = null, modelRoot = '', - targetFolder = '' + targetFolder = '', + useDefaultPaths = null, + useSaveDirAsRoot = false } = {}) { console.warn('[download] downloadVersionWithDefaults: NO fileParams will be sent — backend will always use primary file. ' + 'modelType=%s, modelId=%s, versionId=%s, versionName="%s"', @@ -1824,14 +1828,14 @@ export class DownloadManager { this.modelId = modelId ? modelId.toString() : null; this.source = source; - const useDefaultPaths = !modelRoot; return this.executeDownloadWithProgress({ modelId, versionId, versionName, modelRoot: modelRoot || '', targetFolder: targetFolder || '', - useDefaultPaths, + useDefaultPaths: useDefaultPaths ?? !modelRoot, + useSaveDirAsRoot, source, closeModal: false, }); diff --git a/tests/frontend/components/modelVersionsTab.downloadPath.test.js b/tests/frontend/components/modelVersionsTab.downloadPath.test.js new file mode 100644 index 00000000..98f3ac2c --- /dev/null +++ b/tests/frontend/components/modelVersionsTab.downloadPath.test.js @@ -0,0 +1,210 @@ +import { describe, it, beforeEach, afterEach, expect, vi } from 'vitest'; + +const { + MODEL_VERSIONS_MODULE, + API_FACTORY_MODULE, + DOWNLOAD_MANAGER_MODULE, + UI_HELPERS_MODULE, + STATE_MODULE, + I18N_HELPERS_MODULE, + UTILS_MODULE, +} = vi.hoisted(() => ({ + MODEL_VERSIONS_MODULE: new URL('../../../static/js/components/shared/ModelVersionsTab.js', import.meta.url).pathname, + API_FACTORY_MODULE: new URL('../../../static/js/api/modelApiFactory.js', import.meta.url).pathname, + DOWNLOAD_MANAGER_MODULE: new URL('../../../static/js/managers/DownloadManager.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, + I18N_HELPERS_MODULE: new URL('../../../static/js/utils/i18nHelpers.js', import.meta.url).pathname, + UTILS_MODULE: new URL('../../../static/js/components/shared/utils.js', import.meta.url).pathname, +})); + +const downloadVersionWithDefaults = vi.fn(); + +vi.mock(DOWNLOAD_MANAGER_MODULE, () => ({ + downloadManager: { + downloadVersionWithDefaults, + }, +})); + +vi.mock(UI_HELPERS_MODULE, () => ({ + showToast: vi.fn(), + openCivitaiUrl: vi.fn(), +})); + +const stateMock = { + global: { + settings: { + autoplay_on_hover: false, + version_grouping: 'any', + download_path_templates: { + lora: '{base_model}/{first_tag}', + checkpoint: '{base_model}/{first_tag}', + embedding: '{base_model}/{first_tag}', + }, + }, + }, +}; +vi.mock(STATE_MODULE, () => ({ + state: stateMock, +})); + +vi.mock(I18N_HELPERS_MODULE, () => ({ + translate: vi.fn((_, __, fallback) => fallback ?? ''), +})); + +vi.mock(UTILS_MODULE, () => ({ + formatFileSize: vi.fn(() => '1 MB'), +})); + +vi.mock(API_FACTORY_MODULE, () => ({ + getModelApiClient: vi.fn(), +})); + +const LORA_ROOT = '/models/loras'; + +function buildRecord(targetBaseModel = 'Anima') { + return { + success: true, + record: { + shouldIgnore: false, + inLibraryVersionIds: [10], + versions: [ + { + versionId: 10, + name: 'v1.0', + baseModel: 'Illustrious', + sizeBytes: 1024, + isInLibrary: true, + shouldIgnore: false, + filePath: `${LORA_ROOT}/Illustrious/works/file.safetensors`, + }, + { + versionId: 11, + name: 'v1.1', + baseModel: targetBaseModel, + sizeBytes: 2048, + isInLibrary: false, + shouldIgnore: false, + }, + ], + }, + }; +} + +async function renderAndClickDownload({ currentVersionId = 10, record = null } = {}) { + const { initVersionsTab } = await import(MODEL_VERSIONS_MODULE); + const controller = initVersionsTab({ + modalId: 'model-versions-modal', + modelType: 'loras', + modelId: 123, + currentVersionId, + }); + await controller.load(); + const downloadButton = document.querySelector( + '.model-version-row[data-version-id="11"] [data-version-action="download"]' + ); + downloadButton?.click(); + await new Promise(resolve => setTimeout(resolve, 0)); + return controller; +} + +describe('ModelVersionsTab update download path resolution', () => { + let getModelApiClient; + let fetchModelUpdateVersions; + let fetchModelRoots; + + beforeEach(async () => { + vi.resetModules(); + downloadVersionWithDefaults.mockReset(); + downloadVersionWithDefaults.mockResolvedValue(true); + document.body.innerHTML = ` +
+
+
+
+
+ `; + stateMock.global.settings.version_grouping = 'any'; + stateMock.global.settings.download_path_templates.lora = '{base_model}/{first_tag}'; + ({ getModelApiClient } = await import(API_FACTORY_MODULE)); + fetchModelUpdateVersions = vi.fn(); + fetchModelRoots = vi.fn(); + fetchModelRoots.mockResolvedValue({ roots: [LORA_ROOT] }); + getModelApiClient.mockReturnValue({ + fetchModelUpdateVersions, + fetchModelRoots, + setModelUpdateIgnore: vi.fn(), + setVersionUpdateIgnore: vi.fn(), + deleteModel: vi.fn(), + }); + }); + + afterEach(() => { + document.body.innerHTML = ''; + }); + + it('keeps the current folder when the target version has the same base model', async () => { + fetchModelUpdateVersions.mockResolvedValue(buildRecord('Illustrious')); + + await renderAndClickDownload(); + + expect(downloadVersionWithDefaults).toHaveBeenCalledWith( + 'loras', 123, 11, + expect.objectContaining({ + modelRoot: LORA_ROOT, + targetFolder: 'Illustrious/works', + useDefaultPaths: null, + useSaveDirAsRoot: false, + }) + ); + }); + + it('resolves the template path when the target base model differs and a template is configured', async () => { + fetchModelUpdateVersions.mockResolvedValue(buildRecord()); + + await renderAndClickDownload(); + + expect(downloadVersionWithDefaults).toHaveBeenCalledWith( + 'loras', 123, 11, + expect.objectContaining({ + modelRoot: LORA_ROOT, + targetFolder: '', + useDefaultPaths: true, + useSaveDirAsRoot: true, + }) + ); + }); + + it('keeps the current folder when the target base model differs but no template is configured', async () => { + stateMock.global.settings.download_path_templates.lora = ''; + fetchModelUpdateVersions.mockResolvedValue(buildRecord()); + + await renderAndClickDownload(); + + expect(downloadVersionWithDefaults).toHaveBeenCalledWith( + 'loras', 123, 11, + expect.objectContaining({ + modelRoot: LORA_ROOT, + targetFolder: 'Illustrious/works', + useDefaultPaths: null, + useSaveDirAsRoot: false, + }) + ); + }); + + it('falls back to default paths when no local version exists', async () => { + fetchModelUpdateVersions.mockResolvedValue(buildRecord()); + + await renderAndClickDownload({ currentVersionId: null }); + + expect(downloadVersionWithDefaults).toHaveBeenCalledWith( + 'loras', 123, 11, + expect.objectContaining({ + modelRoot: '', + targetFolder: '', + useDefaultPaths: null, + useSaveDirAsRoot: false, + }) + ); + }); +}); diff --git a/tests/services/test_download_manager_basic.py b/tests/services/test_download_manager_basic.py index 24035b13..2c3a8a01 100644 --- a/tests/services/test_download_manager_basic.py +++ b/tests/services/test_download_manager_basic.py @@ -233,6 +233,58 @@ async def test_successful_download_uses_defaults( assert captured["download_urls"] == ["https://example.invalid/file.safetensors"] +@pytest.mark.asyncio +async def test_download_keeps_save_dir_when_use_save_dir_as_root( + monkeypatch, scanners, metadata_provider, tmp_path +): + """use_default_paths with use_save_dir_as_root resolves the template under + the provided save_dir instead of switching to the default root.""" + manager = DownloadManager() + + captured = {} + + async def fake_execute_download( + self, + *, + download_urls, + save_dir, + metadata, + version_info, + relative_path, + progress_callback, + model_type, + download_id, + transfer_backend=None, + ): + captured.update( + { + "save_dir": Path(save_dir), + "relative_path": relative_path, + "model_type": model_type, + } + ) + return {"success": True} + + monkeypatch.setattr( + DownloadManager, "_execute_download", fake_execute_download, raising=False + ) + + custom_root = tmp_path / "custom_root" + result = await manager.download_from_civitai( + model_version_id=99, + save_dir=str(custom_root), + use_default_paths=True, + use_save_dir_as_root=True, + progress_callback=None, + source=None, + ) + + assert result["success"] is True + assert captured["relative_path"] == "MappedModel/fantasy" + assert captured["save_dir"] == custom_root / "MappedModel" / "fantasy" + assert captured["model_type"] == "lora" + + @pytest.mark.asyncio async def test_successful_download_schedules_auto_example_images( monkeypatch, scanners, metadata_provider, tmp_path @@ -618,6 +670,7 @@ async def test_resume_download_restores_persisted_aria2_task(monkeypatch, tmp_pa use_default_paths=False, source=None, file_params=None, + use_save_dir_as_root=False, ): created.update( { @@ -1037,6 +1090,7 @@ async def test_download_uses_captured_backend_when_settings_change( transfer_backend="python", source=None, file_params=None, + use_save_dir_as_root=False, ): captured["transfer_backend"] = transfer_backend return {"success": True}