mirror of
https://github.com/willmiao/ComfyUI-Lora-Manager.git
synced 2026-08-28 08:21:27 -03:00
fix(download): serialize concurrent downloads resolving to the same target path
This commit is contained in:
@@ -2,6 +2,7 @@
|
|||||||
# Lazy (function-local) imports still count as static edges in basedpyright's
|
# Lazy (function-local) imports still count as static edges in basedpyright's
|
||||||
# reportImportCycles, so the ServiceRegistry singleton pattern necessarily forms
|
# reportImportCycles, so the ServiceRegistry singleton pattern necessarily forms
|
||||||
# import cycles. Breaking them would require an architectural refactor.
|
# import cycles. Breaking them would require an architectural refactor.
|
||||||
|
import contextlib
|
||||||
import copy
|
import copy
|
||||||
import json
|
import json
|
||||||
import logging
|
import logging
|
||||||
@@ -12,6 +13,7 @@ import shutil
|
|||||||
import zipfile
|
import zipfile
|
||||||
from concurrent.futures import ThreadPoolExecutor
|
from concurrent.futures import ThreadPoolExecutor
|
||||||
from collections import OrderedDict
|
from collections import OrderedDict
|
||||||
|
from dataclasses import dataclass, field
|
||||||
import uuid
|
import uuid
|
||||||
from typing import Any, Dict, Iterable, List, Optional, Set, Tuple, cast
|
from typing import Any, Dict, Iterable, List, Optional, Set, Tuple, cast
|
||||||
from urllib.parse import urlparse
|
from urllib.parse import urlparse
|
||||||
@@ -53,6 +55,12 @@ CIVITAI_DOWNLOAD_URL_PREFIXES = (
|
|||||||
NON_DOWNLOADABLE_PRIMARY_TYPES = ("Config", "Archive", "Workflow", "Training Data")
|
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:
|
class DownloadManager:
|
||||||
_instance = None
|
_instance = None
|
||||||
_lock = asyncio.Lock()
|
_lock = asyncio.Lock()
|
||||||
@@ -82,6 +90,11 @@ class DownloadManager:
|
|||||||
self._aria2_state_store = Aria2TransferStateStore()
|
self._aria2_state_store = Aria2TransferStateStore()
|
||||||
self._restored_persisted_downloads = False
|
self._restored_persisted_downloads = False
|
||||||
self._restore_lock = asyncio.Lock()
|
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
|
@staticmethod
|
||||||
def _get_model_download_backend() -> str:
|
def _get_model_download_backend() -> str:
|
||||||
@@ -2132,6 +2145,28 @@ class DownloadManager:
|
|||||||
|
|
||||||
return formatted_path
|
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(
|
async def _execute_download(
|
||||||
self,
|
self,
|
||||||
download_urls: List[str],
|
download_urls: List[str],
|
||||||
@@ -2143,6 +2178,33 @@ class DownloadManager:
|
|||||||
model_type: str = "lora",
|
model_type: str = "lora",
|
||||||
download_id: str | None = None,
|
download_id: str | None = None,
|
||||||
transfer_backend: Optional[str] = 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]:
|
) -> Dict[str, Any]:
|
||||||
"""Execute the actual download process including preview images and model files"""
|
"""Execute the actual download process including preview images and model files"""
|
||||||
metadata_entries: List[Any] = []
|
metadata_entries: List[Any] = []
|
||||||
|
|||||||
@@ -1611,10 +1611,10 @@ async def test_resume_download_requests_reconnect_for_stalled_stream():
|
|||||||
assert pause_control.is_set() is True
|
assert pause_control.is_set() is True
|
||||||
assert pause_control.has_reconnect_request() is True
|
assert pause_control.has_reconnect_request() is True
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_resume_download_rejects_when_not_paused():
|
async def test_resume_download_rejects_when_not_paused():
|
||||||
"""Test that resume_download rejects when download is not paused."""
|
"""Test that resume_download rejects when download is not paused."""
|
||||||
|
|
||||||
manager = DownloadManager()
|
manager = DownloadManager()
|
||||||
|
|
||||||
download_id = "dl"
|
download_id = "dl"
|
||||||
@@ -1624,3 +1624,169 @@ async def test_resume_download_rejects_when_not_paused():
|
|||||||
result = await manager.resume_download(download_id)
|
result = await manager.resume_download(download_id)
|
||||||
|
|
||||||
assert result == {"success": False, "error": "Download is not paused"}
|
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}"
|
||||||
|
)
|
||||||
|
|||||||
Reference in New Issue
Block a user