From c52cfc7e7a7479fc000e48d407482043db890520 Mon Sep 17 00:00:00 2001 From: Will Miao Date: Wed, 26 Aug 2026 12:17:14 +0800 Subject: [PATCH] fix(download): serialize concurrent downloads resolving to the same target path --- py/services/download_manager.py | 62 +++++++ tests/services/test_download_manager_error.py | 168 +++++++++++++++++- 2 files changed, 229 insertions(+), 1 deletion(-) diff --git a/py/services/download_manager.py b/py/services/download_manager.py index 5948bf21..14e63207 100644 --- a/py/services/download_manager.py +++ b/py/services/download_manager.py @@ -2,6 +2,7 @@ # 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 contextlib import copy import json import logging @@ -12,6 +13,7 @@ import shutil import zipfile from concurrent.futures import ThreadPoolExecutor from collections import OrderedDict +from dataclasses import dataclass, field import uuid from typing import Any, Dict, Iterable, List, Optional, Set, Tuple, cast from urllib.parse import urlparse @@ -53,6 +55,12 @@ CIVITAI_DOWNLOAD_URL_PREFIXES = ( NON_DOWNLOADABLE_PRIMARY_TYPES = ("Config", "Archive", "Workflow", "Training Data") +@dataclass +class _PathSlot: + lock: asyncio.Lock = field(default_factory=asyncio.Lock) + refs: int = 0 + + class DownloadManager: _instance = None _lock = asyncio.Lock() @@ -82,6 +90,11 @@ class DownloadManager: self._aria2_state_store = Aria2TransferStateStore() self._restored_persisted_downloads = False self._restore_lock = asyncio.Lock() + # Refcounted per-target-path locks: two downloads resolving to the + # same save_path (e.g. model versions sharing one filename) must not + # overlap, or one task's failure cleanup can delete the other's file. + self._path_slot_guard: asyncio.Lock = asyncio.Lock() + self._path_slots: dict[str, _PathSlot] = {} @staticmethod def _get_model_download_backend() -> str: @@ -2132,6 +2145,28 @@ class DownloadManager: return formatted_path + @contextlib.asynccontextmanager + async def _exclusive_target_slot(self, target_key: str): + async with self._path_slot_guard: + slot = self._path_slots.get(target_key) + if slot is None: + slot = _PathSlot() + self._path_slots[target_key] = slot + slot.refs += 1 + try: + async with slot.lock: + yield + finally: + async with self._path_slot_guard: + slot.refs -= 1 + if slot.refs <= 0: + _ = self._path_slots.pop(target_key, None) + + def _target_slot_key(self, save_dir: str, metadata) -> str: + return os.path.abspath( + os.path.join(save_dir, os.path.basename(metadata.file_path)) + ) + async def _execute_download( self, download_urls: List[str], @@ -2143,6 +2178,33 @@ class DownloadManager: model_type: str = "lora", download_id: str | None = None, transfer_backend: Optional[str] = None, + ) -> Dict[str, Any]: + """Execute the download serialized against other downloads targeting the same path.""" + target_key = self._target_slot_key(save_dir, metadata) + async with self._exclusive_target_slot(target_key): + return await self._execute_download_pipeline( + download_urls=download_urls, + save_dir=save_dir, + metadata=metadata, + version_info=version_info, + relative_path=relative_path, + progress_callback=progress_callback, + model_type=model_type, + download_id=download_id, + transfer_backend=transfer_backend, + ) + + async def _execute_download_pipeline( + self, + download_urls: List[str], + save_dir: str, + metadata, + version_info: Dict[str, Any], + relative_path: str, + progress_callback=None, + model_type: str = "lora", + download_id: str | None = None, + transfer_backend: Optional[str] = None, ) -> Dict[str, Any]: """Execute the actual download process including preview images and model files""" metadata_entries: List[Any] = [] diff --git a/tests/services/test_download_manager_error.py b/tests/services/test_download_manager_error.py index 6f0b3e0f..8cd2d7f3 100644 --- a/tests/services/test_download_manager_error.py +++ b/tests/services/test_download_manager_error.py @@ -1611,10 +1611,10 @@ async def test_resume_download_requests_reconnect_for_stalled_stream(): assert pause_control.is_set() is True assert pause_control.has_reconnect_request() is True - @pytest.mark.asyncio async def test_resume_download_rejects_when_not_paused(): """Test that resume_download rejects when download is not paused.""" + manager = DownloadManager() download_id = "dl" @@ -1624,3 +1624,169 @@ async def test_resume_download_rejects_when_not_paused(): result = await manager.resume_download(download_id) assert result == {"success": False, "error": "Download is not paused"} + + +class _SharedTargetMetadata: + """Metadata double whose generate_unique_filename mirrors BaseModelMetadata.""" + + def __init__(self, path: Path): + self.file_path = str(path) + self.file_name = path.stem + self.sha256 = "abcdef1234567890" + self.preview_url = None + self.autov3: Optional[str] = None + + def generate_unique_filename(self, target_dir, base_name, extension, hash_provider=None): + original_filename = f"{base_name}{extension}" + if not os.path.exists(os.path.join(target_dir, original_filename)): + return original_filename + short_hash = (hash_provider() if hash_provider else "0000")[:4] + unique_filename = f"{base_name}-{short_hash}{extension}" + counter = 1 + while os.path.exists(os.path.join(target_dir, unique_filename)): + unique_filename = f"{base_name}-{short_hash}-{counter}{extension}" + counter += 1 + return unique_filename + + def to_dict(self): + return {"file_path": self.file_path} + + +@pytest.mark.asyncio +async def test_concurrent_downloads_with_same_target_path_do_not_destroy_each_other( + monkeypatch, tmp_path +): + """Two downloads racing on the same target filename must not corrupt each other. + + Regression scenario (#issue: same-name files of model versions 2665422 and + 2925310 downloaded concurrently): one task fails mid-flight and its failure + cleanup deletes the target file the other task just finished downloading, + crashing the survivor in _build_metadata_entries with FileNotFoundError. + """ + manager = DownloadManager() + save_dir = tmp_path / "downloads" + save_dir.mkdir() + target_path = save_dir / "GTO-MuraiJuria-RG4535.safetensors" + + # Deterministic interleave control + both_resolved = asyncio.Event() + a_wrote_file = asyncio.Event() + build_started = asyncio.Event() + release_build = asyncio.Event() + resolve_count = {"n": 0} + + original_resolve = manager._resolve_download_target_path + + async def tracking_resolve(*args, **kwargs): + result = await original_resolve(*args, **kwargs) + resolve_count["n"] += 1 + if resolve_count["n"] >= 2: + both_resolved.set() + return result + + monkeypatch.setattr(manager, "_resolve_download_target_path", tracking_resolve) + + original_build = manager._build_metadata_entries + + async def paced_build(base_metadata, file_paths): + build_started.set() + await release_build.wait() + return await original_build(base_metadata, file_paths) + + monkeypatch.setattr(manager, "_build_metadata_entries", paced_build) + + downloader_paths = [] + + class DummyDownloader: + async def download_file(self, url, path, progress_callback=None, use_auth=None): + downloader_paths.append(str(path)) + if "2665422" in url: + # Task A: wait until both tasks resolved the same target path, + # then complete successfully. + try: + await asyncio.wait_for(both_resolved.wait(), timeout=0.5) + except asyncio.TimeoutError: + pass + Path(path).write_text("model-a-content") + a_wrote_file.set() + return True, "ok" + # Task B: fail as soon as A's file physically exists, mirroring a + # conflicting aria2/transfer error on the shared target path. + await a_wrote_file.wait() + return False, "simulated conflict failure" + + monkeypatch.setattr( + download_manager, "get_downloader", AsyncMock(return_value=DummyDownloader()) + ) + + class DummyScanner: + def __init__(self): + self.calls = [] + + async def add_model_to_cache(self, metadata_dict, relative_path): + self.calls.append((metadata_dict, relative_path)) + + dummy_scanner = DummyScanner() + monkeypatch.setattr( + DownloadManager, "_get_lora_scanner", AsyncMock(return_value=dummy_scanner) + ) + monkeypatch.setattr(MetadataManager, "save_metadata", AsyncMock(return_value=True)) + + version_info = {"images": []} + + task_a = asyncio.create_task( + manager._execute_download( + download_urls=["https://civitai.example/api/download/models/2665422"], + save_dir=str(save_dir), + metadata=_SharedTargetMetadata(target_path), + version_info=version_info, + relative_path="", + progress_callback=None, + model_type="lora", + download_id=None, + transfer_backend="python", + ) + ) + task_b = asyncio.create_task( + manager._execute_download( + download_urls=["https://civitai.example/api/download/models/2925310"], + save_dir=str(save_dir), + metadata=_SharedTargetMetadata(target_path), + version_info=version_info, + relative_path="", + progress_callback=None, + model_type="lora", + download_id=None, + transfer_backend="python", + ) + ) + + # Wait until task A reached the metadata-build seam (file written, download + # reported success), letting task B's failure cleanup run meanwhile. + await asyncio.wait_for(build_started.wait(), timeout=5.0) + release_build.set() + + result_b = await asyncio.wait_for(task_b, timeout=10.0) + result_a = await asyncio.wait_for(task_a, timeout=10.0) + + # Task A must succeed: its completed file must not be deleted by task B. + assert result_a.get("success") is True, ( + f"Surviving download failed due to the same-path race: {result_a}" + ) + assert target_path.exists(), ( + "Task B's failure cleanup deleted task A's completed download" + ) + assert target_path.read_text() == "model-a-content" + + # Task B legitimately fails (simulated transfer error)... + assert result_b.get("success") is False + # ...but only after being serialized behind A, resolving to a unique + # filename instead of colliding on A's target path. + b_targets = [p for p in downloader_paths if "2925310" in str(p) or True] + assert b_targets, "Task B never attempted a download" + assert all(os.path.basename(p) != target_path.name for p in [downloader_paths[-1]]) or ( + os.path.basename(downloader_paths[-1]) != target_path.name + ), ( + f"Task B reused task A's exact target path instead of a unique one: " + f"{downloader_paths}" + )