Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
143 changes: 109 additions & 34 deletions backend/app/runtime/model_downloader.py
Original file line number Diff line number Diff line change
Expand Up @@ -5,6 +5,7 @@
import time
import asyncio
import uuid
import hashlib
import httpx
from pathlib import Path
from typing import Dict, List, Optional
Expand Down Expand Up @@ -68,35 +69,17 @@ def to_info(self) -> DownloadTaskInfo:
error_message=self.error_message,
)

async def run(self):
self.status = "downloading"
self.target_path.parent.mkdir(parents=True, exist_ok=True)

existing_bytes = 0
if self.part_path.exists():
existing_bytes = self.part_path.stat().st_size
self.downloaded_bytes = existing_bytes

headers = {}
if existing_bytes > 0:
headers["Range"] = f"bytes={existing_bytes}-"

async def _download_loop(self, candidate_urls: List[str], headers: dict, existing_bytes: int) -> bool:
"""Download bytes from candidate URLs into part_path."""
last_time = time.monotonic()
bytes_since_last = 0

# Build candidate URL list with selected source_url first, followed by all alternative sources
candidate_urls = [self.source_url] if self.source_url else []
for src in self.model.sources:
if src.url and src.url not in candidate_urls:
candidate_urls.append(src.url)

last_error = None
download_success = False

for url in candidate_urls:
if self._cancel_flag:
self.status = "cancelled"
return
return False

try:
async with httpx.AsyncClient(
Expand All @@ -108,12 +91,43 @@ async def run(self):
if response.status_code in [200, 206]:
content_range = response.headers.get("content-range")
content_length = response.headers.get("content-length")
if content_length and existing_bytes == 0:
self.total_bytes = int(content_length)

mode = "ab" if existing_bytes > 0 and response.status_code == 206 else "wb"
if mode == "wb":
# Validate partial content resumption
if existing_bytes > 0:
if response.status_code == 206:
# Expected Content-Range format: bytes START-END/TOTAL
if content_range:
try:
range_spec = content_range.split()[1]
start_byte = int(range_spec.split("-")[0])
if start_byte != existing_bytes:
# Incompatible range returned by server, reset and redownload from 0
existing_bytes = 0
mode = "wb"
self.downloaded_bytes = 0
else:
mode = "ab"
total_str = range_spec.split("/")[-1]
if total_str != "*":
self.total_bytes = int(total_str)
except Exception:
existing_bytes = 0
mode = "wb"
self.downloaded_bytes = 0
else:
mode = "ab"
else:
# Server returned 200 instead of 206: entire file being sent from 0
mode = "wb"
self.downloaded_bytes = 0
existing_bytes = 0
if content_length:
self.total_bytes = int(content_length)
else:
mode = "wb"
self.downloaded_bytes = 0
if content_length:
self.total_bytes = int(content_length)

with open(self.part_path, mode) as f:
async for chunk in response.aiter_bytes(chunk_size=65536):
Expand All @@ -124,15 +138,15 @@ async def run(self):
self.part_path.unlink()
except Exception:
pass
return
return False

while not self._pause_event.is_set():
self.status = "paused"
self.speed_bps = 0
await asyncio.sleep(0.5)
if self._cancel_flag:
self.status = "cancelled"
return
return False
self.status = "downloading"

f.write(chunk)
Expand Down Expand Up @@ -166,17 +180,78 @@ async def run(self):
self.status = "failed"
self.error_message = f"Failed to download model across candidate sources: {last_error}"
self.speed_bps = 0
return False

return True

async def run(self):
self.status = "downloading"
self.target_path.parent.mkdir(parents=True, exist_ok=True)

existing_bytes = 0
if self.part_path.exists():
existing_bytes = self.part_path.stat().st_size
self.downloaded_bytes = existing_bytes

headers = {}
if existing_bytes > 0:
headers["Range"] = f"bytes={existing_bytes}-"

# Build candidate URL list with selected source_url first, followed by all alternative sources
candidate_urls = [self.source_url] if self.source_url else []
for src in self.model.sources:
if src.url and src.url not in candidate_urls:
candidate_urls.append(src.url)

download_success = await self._download_loop(candidate_urls, headers, existing_bytes)
if not download_success:
return

# Download finished: atomic rename to destination
# Verification: size and checksum checks before promotion
if not self.part_path.exists():
self.status = "failed"
self.error_message = "Downloaded part file does not exist."
self.speed_bps = 0
return

actual_size = self.part_path.stat().st_size

# Check size if total_bytes / model.size_bytes is declared
expected_size = self.total_bytes or self.model.size_bytes
if expected_size and actual_size != expected_size:
self.status = "failed"
self.error_message = (
f"Downloaded artifact size mismatch: expected {expected_size} bytes, "
f"got {actual_size} bytes."
)
self.speed_bps = 0
return

# Check sha256 checksum if provided by catalog or record
if self.model.sha256:
hasher = hashlib.sha256()
with open(self.part_path, "rb") as f:
while chunk := f.read(65536):
hasher.update(chunk)
actual_sha256 = hasher.hexdigest().lower()
expected_sha256 = self.model.sha256.strip().lower()
if actual_sha256 != expected_sha256:
self.status = "failed"
self.error_message = (
f"Downloaded artifact SHA-256 integrity mismatch: expected {expected_sha256}, "
f"got {actual_sha256}."
)
self.speed_bps = 0
return

# Download and integrity verified: atomic replace to destination
# On Windows, os.replace / Path.replace will atomically overwrite target_path without unlinking it first.
# If replace fails, the previous target_path is preserved!
try:
if self.part_path.exists():
if self.target_path.exists():
self.target_path.unlink()
self.part_path.rename(self.target_path)
self.part_path.replace(self.target_path)

self.status = "completed"
self.downloaded_bytes = self.total_bytes
self.downloaded_bytes = actual_size
self.speed_bps = 0
self.eta_seconds = 0

Expand All @@ -197,7 +272,7 @@ async def run(self):
except Exception as e:
if not self._cancel_flag:
self.status = "failed"
self.error_message = str(e)
self.error_message = f"Atomic promotion to destination failed: {e}"
self.speed_bps = 0


Expand Down
132 changes: 132 additions & 0 deletions backend/tests/test_model_downloader.py
Original file line number Diff line number Diff line change
@@ -1,6 +1,7 @@
"""Unit tests for multi-mirror resumable model download engine (MH-M4)."""

import pytest
import asyncio
from pathlib import Path
from unittest.mock import patch, MagicMock

Expand Down Expand Up @@ -113,3 +114,134 @@ def test_downloader_task_lifecycle_controls():
assert downloader.cancel_task("dl_lifecycle_1") is True
assert worker.status == "cancelled"
assert worker._cancel_flag is True


@pytest.mark.asyncio
async def test_download_worker_size_mismatch_rejection(tmp_path):
"""Verify that artifact size mismatch rejects promotion and leaves existing target untouched."""
target_file = tmp_path / "model.safetensors"
target_file.write_text("existing_good_version")
part_file = tmp_path / "model.safetensors.part"
part_file.write_bytes(b"short")

dummy_model = HubModelRecord(
id="test-model",
name="Test Model",
architecture="flux",
category="checkpoint",
version="1.0",
size_bytes=100, # Expects 100 bytes, part is only 5 bytes
parameter_count="12B",
author="test",
description="test",
preview_image_url="test",
min_vram_mb=4000,
optimal_vram_mb=8000,
)

worker = DownloadWorker(
task_id="test_size_fail",
model=dummy_model,
target_engine="comfyui",
target_path=target_file,
source_url="https://example.com/test",
)
# Bypass actual network download loop
worker._download_loop = MagicMock(return_value=asyncio.sleep(0, result=True))

await worker.run()

assert worker.status == "failed"
assert "size mismatch" in worker.error_message
assert target_file.exists()
assert target_file.read_text() == "existing_good_version"
assert part_file.exists()


@pytest.mark.asyncio
async def test_download_worker_sha256_mismatch_rejection(tmp_path):
"""Verify that SHA-256 integrity mismatch rejects promotion and preserves existing target."""
target_file = tmp_path / "model.safetensors"
target_file.write_text("existing_good_version")
part_file = tmp_path / "model.safetensors.part"
part_file.write_bytes(b"corrupted_content")

dummy_model = HubModelRecord(
id="test-model",
name="Test Model",
architecture="flux",
category="checkpoint",
version="1.0",
size_bytes=len(b"corrupted_content"),
sha256="0000000000000000000000000000000000000000000000000000000000000000",
parameter_count="12B",
author="test",
description="test",
preview_image_url="test",
min_vram_mb=4000,
optimal_vram_mb=8000,
)

worker = DownloadWorker(
task_id="test_sha_fail",
model=dummy_model,
target_engine="comfyui",
target_path=target_file,
source_url="https://example.com/test",
)
worker._download_loop = MagicMock(return_value=asyncio.sleep(0, result=True))

await worker.run()

assert worker.status == "failed"
assert "SHA-256 integrity mismatch" in worker.error_message
assert target_file.exists()
assert target_file.read_text() == "existing_good_version"


@pytest.mark.asyncio
async def test_download_worker_atomic_promotion_success(tmp_path):
"""Verify that valid download promotes part file atomically and overwrites destination."""
import hashlib

target_file = tmp_path / "model.safetensors"
target_file.write_text("old_version")
part_file = tmp_path / "model.safetensors.part"
valid_content = b"new_verified_model_weights_12345"
part_file.write_bytes(valid_content)

correct_sha256 = hashlib.sha256(valid_content).hexdigest()

dummy_model = HubModelRecord(
id="test-model",
name="Test Model",
architecture="flux",
category="checkpoint",
version="1.0",
size_bytes=len(valid_content),
sha256=correct_sha256,
parameter_count="12B",
author="test",
description="test",
preview_image_url="test",
min_vram_mb=4000,
optimal_vram_mb=8000,
)

worker = DownloadWorker(
task_id="test_success",
model=dummy_model,
target_engine="comfyui",
target_path=target_file,
source_url="https://example.com/test",
)
worker._download_loop = MagicMock(return_value=asyncio.sleep(0, result=True))

await worker.run()

assert worker.status == "completed"
assert worker.error_message is None
assert not part_file.exists()
assert target_file.exists()
assert target_file.read_bytes() == valid_content

Loading