feat(downloads): return structured 429 rate-limit responses with retry_after

On a CivitAI/CivArchive 429, download-model and download-model-get now
return HTTP 429 with {"reason": "rate_limited", "retry_after": N}
instead of a generic 500 string, and the queue row goes back to
"queued" rather than history as failed — so queue drivers can
auto-pause and retry later instead of burning through the queue.

- new DownloadRateLimitError carrying retry_after/host (opt-in via
  raise_on_rate_limit on Downloader; other call sites keep the legacy
  string behavior)
- fail-fast pre-flight gate in DownloadManager consults
  RateLimitCoordinator before acquiring the semaphore slot: hosts in
  cooldown get an immediate structured 429, no HTTP request attempted
- best-effort 429 detection for the aria2 backend
This commit is contained in:
Will Miao
2026-10-02 21:16:06 +08:00
parent 2a667df98c
commit 034660d8c4
8 changed files with 627 additions and 9 deletions
+4 -2
View File
@@ -1761,7 +1761,8 @@ class ModelDownloadHandler:
payload = await request.json() payload = await request.json()
result = await self._download_use_case.execute(payload) result = await self._download_use_case.execute(payload)
if not result.get("success", False): if not result.get("success", False):
return web.json_response(result, status=500) status = 429 if result.get("reason") == "rate_limited" else 500
return web.json_response(result, status=status)
return web.json_response(result) return web.json_response(result)
except DownloadModelValidationError as exc: except DownloadModelValidationError as exc:
return web.json_response({"success": False, "error": str(exc)}, status=400) return web.json_response({"success": False, "error": str(exc)}, status=400)
@@ -1819,7 +1820,8 @@ class ModelDownloadHandler:
mock_request = type("MockRequest", (), {"json": lambda self=None: future})() mock_request = type("MockRequest", (), {"json": lambda self=None: future})()
result = await self._download_use_case.execute(data) result = await self._download_use_case.execute(data)
if not result.get("success", False): if not result.get("success", False):
return web.json_response(result, status=500) status = 429 if result.get("reason") == "rate_limited" else 500
return web.json_response(result, status=status)
return web.json_response(result) return web.json_response(result)
except DownloadModelValidationError as exc: except DownloadModelValidationError as exc:
return web.json_response({"success": False, "error": str(exc)}, status=400) return web.json_response({"success": False, "error": str(exc)}, status=400)
+135 -1
View File
@@ -44,7 +44,8 @@ from .download_routing import is_diffusion_model_download, resolve_other_downloa
from .settings_manager import get_settings_manager from .settings_manager import get_settings_manager
from .metadata_service import get_default_metadata_provider, get_metadata_provider from .metadata_service import get_default_metadata_provider, get_metadata_provider
from .downloader import get_downloader, DownloadProgress, DownloadStreamControl from .downloader import get_downloader, DownloadProgress, DownloadStreamControl
from .errors import RateLimitError from .errors import DownloadRateLimitError, RateLimitError
from .rate_limit_coordinator import RateLimitCoordinator
from .aria2_downloader import Aria2Error, get_aria2_downloader from .aria2_downloader import Aria2Error, get_aria2_downloader
from .aria2_transfer_state import Aria2TransferStateStore from .aria2_transfer_state import Aria2TransferStateStore
from .download_queue_service import DownloadQueueService from .download_queue_service import DownloadQueueService
@@ -60,6 +61,15 @@ CIVITAI_DOWNLOAD_URL_PREFIXES = (
"https://civitai.red/api/download/", "https://civitai.red/api/download/",
) )
# Hosts a model download may hit (metadata + file transfer). The pre-flight
# cooldown gate consults the RateLimitCoordinator for these before a download
# occupies a concurrency slot.
DOWNLOAD_PREFLIGHT_HOSTS = ("civitai.com", "civitai.red", "civarchive.com")
# Fallback retry_after when neither the vendor nor the coordinator can supply
# a number (matches the Retry-After parsing default in downloader.py).
DEFAULT_RATE_LIMIT_RETRY_AFTER_SECONDS = 60
# File types that are never the intended download target even when CivitAI # File types that are never the intended download target even when CivitAI
# marks them primary — configs/archives/workflows are auxiliary artifacts. # marks them primary — configs/archives/workflows are auxiliary artifacts.
@@ -203,11 +213,30 @@ class DownloadManager:
) )
except Aria2Error as exc: except Aria2Error as exc:
logger.error("aria2 download failed for %s: %s", download_url, exc) logger.error("aria2 download failed for %s: %s", download_url, exc)
# Best-effort 429 detection: aria2 reports HTTP status errors
# via its error message (e.g. "status=429") without exposing
# the vendor's Retry-After. Surface the structured rate-limit
# error so the queue contract behaves the same as the python
# backend; the coordinator backoff supplies the wait time.
message = str(exc)
if "429" in message or "rate limit" in message.lower():
host = self._url_host(download_url)
coordinator = await RateLimitCoordinator.get_instance()
if coordinator.enabled:
coordinator.register_rate_limit(host, None)
raise DownloadRateLimitError(
f"Download rate limited (429): {message}",
retry_after=None,
host=host,
) from exc
return False, str(exc) return False, str(exc)
download_kwargs: Dict[str, Any] = { download_kwargs: Dict[str, Any] = {
"progress_callback": progress_callback, "progress_callback": progress_callback,
"use_auth": use_auth, "use_auth": use_auth,
# The model download queue contract requires structured 429
# propagation (reason="rate_limited"), not a plain error string.
"raise_on_rate_limit": True,
} }
if pause_control is not None: if pause_control is not None:
@@ -216,6 +245,88 @@ class DownloadManager:
downloader = await get_downloader() downloader = await get_downloader()
return await downloader.download_file(download_url, save_path, **download_kwargs) return await downloader.download_file(download_url, save_path, **download_kwargs)
@staticmethod
def _url_host(url: str) -> str:
"""Extract the normalized hostname from a URL (fallback: ``unknown``)."""
hostname = urlparse(url).hostname
return hostname.lower() if hostname else "unknown"
async def _preflight_rate_limit_error(self) -> Optional[DownloadRateLimitError]:
"""Fail fast when a download target host is in a rate-limit cooldown.
Consults the RateLimitCoordinator's per-host cooldown state for the
hosts a model download may hit. Runs BEFORE the concurrency semaphore
is acquired so queued items never occupy a slot during a 429 episode.
Deliberately non-blocking: the caller is expected to pace itself (the
companion extension auto-pauses on the structured 429 response).
"""
coordinator = await RateLimitCoordinator.get_instance()
worst_host: Optional[str] = None
worst_remaining = 0.0
for host in DOWNLOAD_PREFLIGHT_HOSTS:
remaining = coordinator.remaining_seconds(host)
if remaining > worst_remaining:
worst_host = host
worst_remaining = remaining
if worst_host is None or worst_remaining <= 0:
return None
retry_after = max(1, int(worst_remaining + 0.5))
return DownloadRateLimitError(
f"Download rate limited: '{worst_host}' is in cooldown, "
f"retry after {retry_after}s",
retry_after=worst_remaining,
host=worst_host,
)
async def _handle_rate_limited_download(
self,
task_id: str,
exc: RateLimitError,
) -> Dict[str, Any]:
"""Build the structured rate-limit result for a failed download.
The queue row goes back to ``queued`` (NOT history) so a later retry
simply starts the download again — this is what lets the companion
extension auto-pause the queue during a 429 episode and resume it
after ``retry_after`` seconds.
"""
retry_after = exc.retry_after
host = getattr(exc, "host", None) or exc.provider
if retry_after is None or retry_after <= 0:
# The coordinator clamps/backoffs via register_rate_limit, so its
# remaining cooldown supplies the number when the vendor didn't.
coordinator = await RateLimitCoordinator.get_instance()
remaining = coordinator.remaining_seconds(host)
if remaining > 0:
retry_after = remaining
if retry_after is None or retry_after <= 0:
retry_after = float(DEFAULT_RATE_LIMIT_RETRY_AFTER_SECONDS)
retry_after_seconds = max(1, int(retry_after + 0.5))
message = str(exc) or (
f"Download rate limited, retry after {retry_after_seconds}s"
)
if task_id in self._active_downloads:
self._active_downloads[task_id]["status"] = "queued"
self._active_downloads[task_id]["error"] = message
self._active_downloads[task_id]["bytes_per_second"] = 0.0
try:
queue_service = await DownloadQueueService.get_instance()
await queue_service.update_status(task_id, "queued", error=message)
except Exception:
logger.warning(
"Failed to re-queue rate-limited download %s", task_id, exc_info=True
)
return {
"success": False,
"reason": "rate_limited",
"retry_after": retry_after_seconds,
"error": message,
}
async def _get_lora_scanner(self): async def _get_lora_scanner(self):
"""Get the lora scanner from registry""" """Get the lora scanner from registry"""
return await ServiceRegistry.get_lora_scanner() return await ServiceRegistry.get_lora_scanner()
@@ -558,6 +669,15 @@ class DownloadManager:
original_callback, snapshot, progress_value original_callback, snapshot, progress_value
) )
# Pre-flight cooldown gate: fail fast (without holding a semaphore
# slot) when a target host is still cooling down from an earlier 429.
preflight_error = await self._preflight_rate_limit_error()
if preflight_error is not None:
logger.info(
"Download %s skipped: %s", task_id, preflight_error
)
return await self._handle_rate_limited_download(task_id, preflight_error)
# Acquire semaphore to limit concurrent downloads # Acquire semaphore to limit concurrent downloads
try: try:
async with self._download_semaphore: async with self._download_semaphore:
@@ -662,6 +782,12 @@ class DownloadManager:
logger.info(f"Download cancelled for task {task_id}") logger.info(f"Download cancelled for task {task_id}")
raise raise
except RateLimitError as e:
# 429 (real vendor response or cooldown gate): re-queue
# instead of completing as failed so a later retry just
# starts the download again.
logger.info(f"Download rate limited for task {task_id}: {e}")
return await self._handle_rate_limited_download(task_id, e)
except Exception as e: except Exception as e:
# Handle other errors # Handle other errors
logger.error( logger.error(
@@ -2115,6 +2241,10 @@ class DownloadManager:
return result return result
except RateLimitError:
# Structured 429 propagation must reach _download_with_semaphore
# unmodified so the queue row is re-queued instead of failed.
raise
except Exception as e: except Exception as e:
logger.error(f"Error in download_from_civitai: {e}", exc_info=True) logger.error(f"Error in download_from_civitai: {e}", exc_info=True)
# Check if this might be an early access error # Check if this might be an early access error
@@ -2837,6 +2967,10 @@ class DownloadManager:
return {"success": True} return {"success": True}
except RateLimitError:
# Structured 429 propagation must reach _download_with_semaphore
# unmodified so the queue row is re-queued instead of failed.
raise
except Exception as e: except Exception as e:
logger.error(f"Error in _execute_download: {e}", exc_info=True) logger.error(f"Error in _execute_download: {e}", exc_info=True)
cleanup_targets = { cleanup_targets = {
+33 -1
View File
@@ -31,7 +31,7 @@ from .connectivity_guard import (
OFFLINE_FRIENDLY_MESSAGE, OFFLINE_FRIENDLY_MESSAGE,
ConnectivityGuard, ConnectivityGuard,
) )
from .errors import RateLimitError from .errors import DownloadRateLimitError, RateLimitError
from .rate_limit_coordinator import RateLimitCoordinator from .rate_limit_coordinator import RateLimitCoordinator
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
@@ -434,6 +434,7 @@ class Downloader:
custom_headers: Optional[Dict[str, str]] = None, custom_headers: Optional[Dict[str, str]] = None,
allow_resume: bool = True, allow_resume: bool = True,
pause_event: Optional[DownloadStreamControl] = None, pause_event: Optional[DownloadStreamControl] = None,
raise_on_rate_limit: bool = False,
) -> Tuple[bool, str]: ) -> Tuple[bool, str]:
""" """
Download a file with resumable downloads and retry mechanism Download a file with resumable downloads and retry mechanism
@@ -446,6 +447,11 @@ class Downloader:
custom_headers: Additional headers to include in request custom_headers: Additional headers to include in request
allow_resume: Whether to support resumable downloads allow_resume: Whether to support resumable downloads
pause_event: Optional stream control used to pause/resume and request reconnects pause_event: Optional stream control used to pause/resume and request reconnects
raise_on_rate_limit: When True, a 429 response raises
``DownloadRateLimitError`` instead of returning a plain error
string, so callers (the model download manager) can build the
structured rate-limit result required by the download queue
contract. Defaults to the legacy tuple behavior.
Returns: Returns:
Tuple[bool, str]: (success, save_path or error message) Tuple[bool, str]: (success, save_path or error message)
@@ -610,6 +616,12 @@ class Downloader:
logger.warning( logger.warning(
f"Rate limited (429) for {url}, retry_after={retry_after}" f"Rate limited (429) for {url}, retry_after={retry_after}"
) )
if raise_on_rate_limit:
raise DownloadRateLimitError(
f"Download rate limited (429), retry after {retry_after}s",
retry_after=retry_after,
host=self._guard_destination(url),
)
return False, f"Download rate limited (429), retry after {retry_after}s" return False, f"Download rate limited (429), retry after {retry_after}s"
else: else:
logger.error( logger.error(
@@ -902,6 +914,11 @@ class Downloader:
f"Network error after {self.max_retries + 1} attempts: {str(e)}", f"Network error after {self.max_retries + 1} attempts: {str(e)}",
) )
except DownloadRateLimitError:
# 429s are never retried in-band; the structured error must
# reach the caller (download manager) unmodified.
raise
except Exception as e: except Exception as e:
logger.error(f"Unexpected download error: {e}") logger.error(f"Unexpected download error: {e}")
return False, str(e) return False, str(e)
@@ -931,6 +948,7 @@ class Downloader:
use_auth: bool = False, use_auth: bool = False,
custom_headers: Optional[Dict[str, str]] = None, custom_headers: Optional[Dict[str, str]] = None,
return_headers: bool = False, return_headers: bool = False,
raise_on_rate_limit: bool = False,
) -> Tuple[bool, Union[bytes, str], Optional[Dict[str, Any]]]: ) -> Tuple[bool, Union[bytes, str], Optional[Dict[str, Any]]]:
""" """
Download a file to memory (for small files like preview images) Download a file to memory (for small files like preview images)
@@ -940,6 +958,10 @@ class Downloader:
use_auth: Whether to include authentication headers use_auth: Whether to include authentication headers
custom_headers: Additional headers to include in request custom_headers: Additional headers to include in request
return_headers: Whether to return response headers along with content return_headers: Whether to return response headers along with content
raise_on_rate_limit: When True, a 429 response raises
``DownloadRateLimitError`` instead of returning a plain error
string (see ``download_file``). Defaults to the legacy tuple
behavior.
Returns: Returns:
Tuple[bool, Union[bytes, str], Optional[Dict]]: (success, content or error message, response headers if requested) Tuple[bool, Union[bytes, str], Optional[Dict]]: (success, content or error message, response headers if requested)
@@ -1002,11 +1024,21 @@ class Downloader:
"Rate limited (429) for %s, no Retry-After header; defaulting to %ss", "Rate limited (429) for %s, no Retry-After header; defaulting to %ss",
url, retry_after, url, retry_after,
) )
if raise_on_rate_limit:
raise DownloadRateLimitError(
f"Rate limited (429), retry after {retry_after}s",
retry_after=retry_after,
host=destination,
)
return False, f"Rate limited (429), retry after {retry_after}s", None return False, f"Rate limited (429), retry after {retry_after}s", None
else: else:
error_msg = f"Download failed with status {response.status}" error_msg = f"Download failed with status {response.status}"
return False, error_msg, None return False, error_msg, None
except DownloadRateLimitError:
# Structured rate-limit errors must reach the caller unmodified.
raise
except Exception as e: except Exception as e:
if guard.is_network_unreachable_error(e): if guard.is_network_unreachable_error(e):
guard.register_network_failure(e, destination) guard.register_network_failure(e, destination)
+19
View File
@@ -20,6 +20,25 @@ class RateLimitError(RuntimeError):
self.provider = provider self.provider = provider
class DownloadRateLimitError(RateLimitError):
"""Raised when a file download is rejected with HTTP 429.
Carries the vendor's ``Retry-After`` hint (when present) and the target
host so the download manager can build the structured rate-limit result
the download queue contract expects.
"""
def __init__(
self,
message: str,
*,
retry_after: Optional[float] = None,
host: Optional[str] = None,
) -> None:
super().__init__(message, retry_after=retry_after)
self.host = host
class ResourceNotFoundError(RuntimeError): class ResourceNotFoundError(RuntimeError):
"""Raised when a remote resource is permanently missing.""" """Raised when a remote resource is permanently missing."""
@@ -9,6 +9,8 @@ with ``success: false``, not 404. The browser extension's apiFetch treats any
import json import json
import logging import logging
from pathlib import Path from pathlib import Path
from types import SimpleNamespace
from unittest.mock import AsyncMock
import pytest import pytest
from aiohttp import web from aiohttp import web
@@ -184,3 +186,83 @@ async def test_retry_failed_history_returns_success(
queue = await queue_service.get_queue() queue = await queue_service.get_queue()
assert len(queue) == 1 assert len(queue) == 1
assert queue[0]["status"] == "queued" assert queue[0]["status"] == "queued"
# ----------------------------------------------------------------------
# Structured 429 rate-limit responses (download queue contract)
# ----------------------------------------------------------------------
_RATE_LIMITED_RESULT = {
"success": False,
"reason": "rate_limited",
"retry_after": 120,
"error": "Download rate limited (429), retry after 120s",
"download_id": "dl-1",
}
def _make_download_handler(result: dict) -> ModelDownloadHandler:
return ModelDownloadHandler(
ws_manager=None, # pyright: ignore[reportArgumentType] - unused by download endpoints
logger=logging.getLogger("test-download-rate-limit"),
download_use_case=SimpleNamespace(execute=AsyncMock(return_value=result)), # pyright: ignore[reportArgumentType]
download_coordinator=None, # pyright: ignore[reportArgumentType] - unused by download endpoints
)
@pytest.mark.asyncio
async def test_download_model_post_returns_429_when_rate_limited() -> None:
"""POST /download-model maps reason="rate_limited" to HTTP 429."""
handler = _make_download_handler(dict(_RATE_LIMITED_RESULT))
request = SimpleNamespace(json=AsyncMock(return_value={"model_id": 1}))
response = await handler.download_model(request)
assert response.status == 429
payload = json.loads(response.text)
assert payload["success"] is False
assert payload["reason"] == "rate_limited"
assert payload["retry_after"] == 120
assert payload["error"]
@pytest.mark.asyncio
async def test_download_model_post_keeps_500_for_other_failures() -> None:
handler = _make_download_handler({"success": False, "error": "boom"})
request = SimpleNamespace(json=AsyncMock(return_value={"model_id": 1}))
response = await handler.download_model(request)
assert response.status == 500
payload = json.loads(response.text)
assert payload["success"] is False
assert "reason" not in payload
@pytest.mark.asyncio
async def test_download_model_get_returns_429_when_rate_limited() -> None:
"""GET /download-model-get (extension queue driver) gets the same 429."""
handler = _make_download_handler(dict(_RATE_LIMITED_RESULT))
request = _queue_request("/api/lm/download-model-get", {"model_id": "1"})
response = await handler.download_model_get(request)
assert response.status == 429
payload = json.loads(response.text)
assert payload["success"] is False
assert payload["reason"] == "rate_limited"
assert payload["retry_after"] == 120
assert payload["error"]
@pytest.mark.asyncio
async def test_download_model_get_keeps_500_for_other_failures() -> None:
handler = _make_download_handler({"success": False, "error": "boom"})
request = _queue_request("/api/lm/download-model-get", {"model_id": "1"})
response = await handler.download_model_get(request)
assert response.status == 500
payload = json.loads(response.text)
assert payload["success"] is False
assert "reason" not in payload
@@ -130,7 +130,7 @@ async def test_execute_download_uses_rewritten_civitai_preview(monkeypatch, tmp_
self.file_calls: list[tuple[str, str]] = [] self.file_calls: list[tuple[str, str]] = []
self.memory_calls = 0 self.memory_calls = 0
async def download_file(self, url, path, progress_callback=None, use_auth=None): async def download_file(self, url, path, progress_callback=None, use_auth=None, **_kwargs):
self.file_calls.append((url, path)) self.file_calls.append((url, path))
if url.endswith(".jpeg"): if url.endswith(".jpeg"):
Path(path).write_bytes(b"preview") Path(path).write_bytes(b"preview")
@@ -248,7 +248,7 @@ async def test_execute_download_respects_blur_setting(monkeypatch, tmp_path):
def __init__(self): def __init__(self):
self.file_calls: list[tuple[str, str]] = [] self.file_calls: list[tuple[str, str]] = []
async def download_file(self, url, path, progress_callback=None, use_auth=None): async def download_file(self, url, path, progress_callback=None, use_auth=None, **_kwargs):
self.file_calls.append((url, path)) self.file_calls.append((url, path))
if url.endswith(".safetensors"): if url.endswith(".safetensors"):
Path(path).write_bytes(b"model") Path(path).write_bytes(b"model")
+237 -3
View File
@@ -12,9 +12,12 @@ from unittest.mock import AsyncMock
import pytest import pytest
from py.services.download_manager import DownloadManager from py.services.download_manager import DownloadManager
from py.services.download_queue_service import DownloadQueueService
from py.services.downloader import DownloadStreamControl from py.services.downloader import DownloadStreamControl
from py.services import download_manager from py.services import download_manager
from py.services import aria2_transfer_state from py.services import aria2_transfer_state
from py.services.errors import DownloadRateLimitError, RateLimitError
from py.services.rate_limit_coordinator import RateLimitCoordinator
from py.services.service_registry import ServiceRegistry from py.services.service_registry import ServiceRegistry
from py.services.settings_manager import SettingsManager, get_settings_manager from py.services.settings_manager import SettingsManager, get_settings_manager
from py.utils.metadata_manager import MetadataManager from py.utils.metadata_manager import MetadataManager
@@ -60,6 +63,25 @@ def isolate_aria2_state(monkeypatch, tmp_path):
) )
@pytest.fixture
def queue_service(tmp_path, monkeypatch):
"""Return a tmp-backed DownloadQueueService and stub the singleton."""
service = DownloadQueueService(db_path=str(tmp_path / "queue.sqlite"))
async def fake_get_instance(_cls=None):
return service
monkeypatch.setattr(DownloadQueueService, "get_instance", fake_get_instance)
return service
@pytest.fixture
def reset_rate_limit_coordinator():
RateLimitCoordinator._instance = None
yield
RateLimitCoordinator._instance = None
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_execute_download_retries_urls(monkeypatch, tmp_path): async def test_execute_download_retries_urls(monkeypatch, tmp_path):
"""Test that download retries multiple URLs on failure.""" """Test that download retries multiple URLs on failure."""
@@ -97,7 +119,9 @@ async def test_execute_download_retries_urls(monkeypatch, tmp_path):
def __init__(self): def __init__(self):
self.calls = [] self.calls = []
async def download_file(self, url, path, progress_callback=None, use_auth=None): async def download_file(
self, url, path, progress_callback=None, use_auth=None, **_kwargs
):
self.calls.append((url, path, use_auth)) self.calls.append((url, path, use_auth))
if len(self.calls) == 1: if len(self.calls) == 1:
return False, "first failed" return False, "first failed"
@@ -373,7 +397,7 @@ async def test_execute_download_adjusts_checkpoint_sub_type(monkeypatch, tmp_pat
class DummyDownloader: class DummyDownloader:
async def download_file( async def download_file(
self, _url, path, progress_callback=None, use_auth=None self, _url, path, progress_callback=None, use_auth=None, **_kwargs
): ):
Path(path).write_text("content") Path(path).write_text("content")
return True, "ok" return True, "ok"
@@ -1698,7 +1722,9 @@ async def test_concurrent_downloads_with_same_target_path_do_not_destroy_each_ot
downloader_paths = [] downloader_paths = []
class DummyDownloader: class DummyDownloader:
async def download_file(self, url, path, progress_callback=None, use_auth=None): async def download_file(
self, url, path, progress_callback=None, use_auth=None, **_kwargs
):
downloader_paths.append(str(path)) downloader_paths.append(str(path))
if "2665422" in url: if "2665422" in url:
# Task A: wait until both tasks resolved the same target path, # Task A: wait until both tasks resolved the same target path,
@@ -1927,3 +1953,211 @@ async def test_restore_persisted_downloads_removes_orphaned_aria2_control_file(
assert not control_path.exists() assert not control_path.exists()
assert await manager._aria2_state_store.get(download_id) is None assert await manager._aria2_state_store.get(download_id) is None
# ----------------------------------------------------------------------
# Structured 429 rate-limit handling (download queue contract)
# ----------------------------------------------------------------------
def _prepare_tracked_download(manager, download_id, status="downloading"):
manager._active_downloads[download_id] = {
"transfer_backend": "python",
"status": status,
"bytes_per_second": 0.0,
}
manager._pause_events[download_id] = DownloadStreamControl()
async def _run_download(manager, download_id, tmp_path):
return await manager._download_with_semaphore(
download_id,
1,
None,
str(tmp_path),
"",
None,
False,
None,
None,
False,
)
@pytest.mark.asyncio
async def test_rate_limited_download_is_requeued_not_failed(
monkeypatch, tmp_path, queue_service, reset_rate_limit_coordinator
):
"""A 429 during the transfer re-queues the item with a structured result."""
manager = DownloadManager()
monkeypatch.setattr(manager, "_cleanup_download_record", AsyncMock())
download_id = "dl-429"
await queue_service.add_to_queue(download_id=download_id, model_id=1)
await queue_service.update_status(download_id, "downloading")
_prepare_tracked_download(manager, download_id)
monkeypatch.setattr(
manager,
"_execute_original_download",
AsyncMock(
side_effect=DownloadRateLimitError(
"Download rate limited (429), retry after 120s",
retry_after=120,
host="civitai.com",
)
),
)
result = await _run_download(manager, download_id, tmp_path)
assert result["success"] is False
assert result["reason"] == "rate_limited"
assert result["retry_after"] == 120
assert "rate limited" in result["error"].lower()
# Back to "queued" — NOT moved to history as failed.
queue = await queue_service.get_queue()
assert len(queue) == 1
assert queue[0]["status"] == "queued"
history = await queue_service.get_history()
assert history["items"] == []
assert manager._active_downloads[download_id]["status"] == "queued"
@pytest.mark.asyncio
async def test_preflight_gate_blocks_when_host_in_cooldown(
monkeypatch, tmp_path, queue_service, reset_rate_limit_coordinator
):
"""Host in cooldown: immediate structured 429, no transfer, no slot held."""
manager = DownloadManager()
monkeypatch.setattr(manager, "_cleanup_download_record", AsyncMock())
coordinator = await RateLimitCoordinator.get_instance()
coordinator.register_rate_limit("civitai.com", 300)
download_id = "dl-cooldown"
await queue_service.add_to_queue(download_id=download_id, model_id=1)
_prepare_tracked_download(manager, download_id, status="waiting")
execute = AsyncMock(
side_effect=AssertionError("download must not start during cooldown")
)
monkeypatch.setattr(manager, "_execute_original_download", execute)
result = await _run_download(manager, download_id, tmp_path)
assert execute.await_count == 0
assert result["success"] is False
assert result["reason"] == "rate_limited"
assert 290 <= result["retry_after"] <= 300
queue = await queue_service.get_queue()
assert len(queue) == 1
assert queue[0]["status"] == "queued"
history = await queue_service.get_history()
assert history["items"] == []
# The concurrency slot was never occupied by the gated download.
await asyncio.wait_for(manager._download_semaphore.acquire(), timeout=0.1)
manager._download_semaphore.release()
@pytest.mark.asyncio
async def test_rate_limit_retry_after_falls_back_to_coordinator_backoff(
monkeypatch, tmp_path, queue_service, reset_rate_limit_coordinator
):
"""When the 429 carries no Retry-After, the coordinator backoff supplies it."""
manager = DownloadManager()
monkeypatch.setattr(manager, "_cleanup_download_record", AsyncMock())
# Cooldown on a mirror host that is NOT in the pre-flight host list, so
# the failure must come from the transfer itself.
coordinator = await RateLimitCoordinator.get_instance()
backoff = coordinator.register_rate_limit("mirror.example.com", None)
assert backoff == 30.0
download_id = "dl-429-mirror"
await queue_service.add_to_queue(download_id=download_id, model_id=1)
_prepare_tracked_download(manager, download_id)
monkeypatch.setattr(
manager,
"_execute_original_download",
AsyncMock(
side_effect=DownloadRateLimitError(
"Download rate limited (429)",
retry_after=None,
host="mirror.example.com",
)
),
)
result = await _run_download(manager, download_id, tmp_path)
assert result["success"] is False
assert result["reason"] == "rate_limited"
assert 25 <= result["retry_after"] <= 30
queue = await queue_service.get_queue()
assert queue[0]["status"] == "queued"
@pytest.mark.asyncio
async def test_metadata_fetch_rate_limit_error_is_also_requeued(
monkeypatch, tmp_path, queue_service, reset_rate_limit_coordinator
):
"""A plain RateLimitError (metadata fetch path) gets the same treatment."""
manager = DownloadManager()
monkeypatch.setattr(manager, "_cleanup_download_record", AsyncMock())
download_id = "dl-meta-429"
await queue_service.add_to_queue(download_id=download_id, model_id=1)
_prepare_tracked_download(manager, download_id)
monkeypatch.setattr(
manager,
"_execute_original_download",
AsyncMock(
side_effect=RateLimitError(
"Request rate limited", retry_after=45, provider="civitai_api"
)
),
)
result = await _run_download(manager, download_id, tmp_path)
assert result["success"] is False
assert result["reason"] == "rate_limited"
assert result["retry_after"] == 45
queue = await queue_service.get_queue()
assert queue[0]["status"] == "queued"
history = await queue_service.get_history()
assert history["items"] == []
@pytest.mark.asyncio
async def test_generic_failure_still_moves_to_history_as_failed(
monkeypatch, tmp_path, queue_service, reset_rate_limit_coordinator
):
"""Non-rate-limit failures keep the legacy move-to-history behavior."""
manager = DownloadManager()
monkeypatch.setattr(manager, "_cleanup_download_record", AsyncMock())
download_id = "dl-generic-fail"
await queue_service.add_to_queue(download_id=download_id, model_id=1)
_prepare_tracked_download(manager, download_id)
monkeypatch.setattr(
manager,
"_execute_original_download",
AsyncMock(side_effect=RuntimeError("disk full")),
)
result = await _run_download(manager, download_id, tmp_path)
assert result == {"success": False, "error": "disk full"}
queue = await queue_service.get_queue()
assert queue == []
history = await queue_service.get_history()
assert len(history["items"]) == 1
assert history["items"][0]["status"] == "failed"
+115
View File
@@ -5,7 +5,20 @@ from typing import Sequence
import pytest import pytest
from py.services.connectivity_guard import ConnectivityGuard
from py.services.downloader import Downloader from py.services.downloader import Downloader
from py.services.errors import DownloadRateLimitError
from py.services.rate_limit_coordinator import RateLimitCoordinator
@pytest.fixture(autouse=True)
def _reset_rate_limit_singletons():
"""429 handling mutates the global coordinator/guard; isolate per test."""
RateLimitCoordinator._instance = None
ConnectivityGuard._instance = None
yield
RateLimitCoordinator._instance = None
ConnectivityGuard._instance = None
class FakeStream: class FakeStream:
@@ -297,3 +310,105 @@ async def test_disable_netrc_auth_ignores_netrc_file(tmp_path, monkeypatch):
_disable_netrc_auth(session) _disable_netrc_auth(session)
assert session._get_netrc_auth("civitai.red") is None assert session._get_netrc_auth("civitai.red") is None
@pytest.mark.asyncio
async def test_download_file_rate_limit_raises_structured_error_when_enabled(tmp_path):
"""With raise_on_rate_limit=True a 429 raises DownloadRateLimitError."""
target_path = tmp_path / "model" / "file.bin"
target_path.parent.mkdir()
responses = [
lambda: FakeResponse(
status=429,
headers={"Retry-After": "120"},
chunks=[],
)
]
downloader = _build_downloader(responses)
with pytest.raises(DownloadRateLimitError) as exc_info:
await downloader.download_file(
"https://example.com/file",
str(target_path),
raise_on_rate_limit=True,
)
assert exc_info.value.retry_after == 120.0
assert exc_info.value.host == "example.com"
# The cooldown was registered with the coordinator for the target host.
coordinator = await RateLimitCoordinator.get_instance()
assert coordinator.remaining_seconds("example.com") > 0
# 429 is never retried in-band.
assert _session(downloader)._get_calls == 1
assert not Path(str(target_path) + ".part").exists()
@pytest.mark.asyncio
async def test_download_file_rate_limit_keeps_legacy_tuple_by_default(tmp_path):
"""Default behavior (other callers) stays the plain error string."""
target_path = tmp_path / "model" / "file.bin"
target_path.parent.mkdir()
responses = [
lambda: FakeResponse(
status=429,
headers={"Retry-After": "120"},
chunks=[],
)
]
downloader = _build_downloader(responses)
success, message = await downloader.download_file(
"https://example.com/file", str(target_path)
)
assert success is False
assert message == "Download rate limited (429), retry after 120.0s"
@pytest.mark.asyncio
async def test_download_to_memory_rate_limit_raises_structured_error_when_enabled():
responses = [
lambda: FakeResponse(
status=429,
headers={"Retry-After": "90"},
chunks=[],
)
]
downloader = _build_downloader(responses)
with pytest.raises(DownloadRateLimitError) as exc_info:
await downloader.download_to_memory(
"https://example.com/preview.png",
raise_on_rate_limit=True,
)
assert exc_info.value.retry_after == 90
assert exc_info.value.host == "example.com"
coordinator = await RateLimitCoordinator.get_instance()
assert coordinator.remaining_seconds("example.com") > 0
@pytest.mark.asyncio
async def test_download_to_memory_rate_limit_keeps_legacy_tuple_by_default():
responses = [
lambda: FakeResponse(
status=429,
headers={"Retry-After": "90"},
chunks=[],
)
]
downloader = _build_downloader(responses)
success, message, headers = await downloader.download_to_memory(
"https://example.com/preview.png"
)
assert success is False
assert message == "Rate limited (429), retry after 90s"
assert headers is None