diff --git a/py/routes/handlers/model_handlers.py b/py/routes/handlers/model_handlers.py index 43d7cbc2..a2287ef7 100644 --- a/py/routes/handlers/model_handlers.py +++ b/py/routes/handlers/model_handlers.py @@ -1761,7 +1761,8 @@ class ModelDownloadHandler: payload = await request.json() result = await self._download_use_case.execute(payload) 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) except DownloadModelValidationError as exc: 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})() result = await self._download_use_case.execute(data) 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) except DownloadModelValidationError as exc: return web.json_response({"success": False, "error": str(exc)}, status=400) diff --git a/py/services/download_manager.py b/py/services/download_manager.py index 62fff77a..570f19c9 100644 --- a/py/services/download_manager.py +++ b/py/services/download_manager.py @@ -44,7 +44,8 @@ from .download_routing import is_diffusion_model_download, resolve_other_downloa from .settings_manager import get_settings_manager from .metadata_service import get_default_metadata_provider, get_metadata_provider 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_transfer_state import Aria2TransferStateStore from .download_queue_service import DownloadQueueService @@ -60,6 +61,15 @@ CIVITAI_DOWNLOAD_URL_PREFIXES = ( "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 # marks them primary — configs/archives/workflows are auxiliary artifacts. @@ -203,11 +213,30 @@ class DownloadManager: ) except Aria2Error as 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) download_kwargs: Dict[str, Any] = { "progress_callback": progress_callback, "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: @@ -216,6 +245,88 @@ class DownloadManager: downloader = await get_downloader() 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): """Get the lora scanner from registry""" return await ServiceRegistry.get_lora_scanner() @@ -558,6 +669,15 @@ class DownloadManager: 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 try: async with self._download_semaphore: @@ -662,6 +782,12 @@ class DownloadManager: logger.info(f"Download cancelled for task {task_id}") 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: # Handle other errors logger.error( @@ -2115,6 +2241,10 @@ class DownloadManager: 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: logger.error(f"Error in download_from_civitai: {e}", exc_info=True) # Check if this might be an early access error @@ -2837,6 +2967,10 @@ class DownloadManager: 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: logger.error(f"Error in _execute_download: {e}", exc_info=True) cleanup_targets = { diff --git a/py/services/downloader.py b/py/services/downloader.py index 94be53a7..246bfc77 100644 --- a/py/services/downloader.py +++ b/py/services/downloader.py @@ -31,7 +31,7 @@ from .connectivity_guard import ( OFFLINE_FRIENDLY_MESSAGE, ConnectivityGuard, ) -from .errors import RateLimitError +from .errors import DownloadRateLimitError, RateLimitError from .rate_limit_coordinator import RateLimitCoordinator logger = logging.getLogger(__name__) @@ -434,6 +434,7 @@ class Downloader: custom_headers: Optional[Dict[str, str]] = None, allow_resume: bool = True, pause_event: Optional[DownloadStreamControl] = None, + raise_on_rate_limit: bool = False, ) -> Tuple[bool, str]: """ Download a file with resumable downloads and retry mechanism @@ -446,6 +447,11 @@ class Downloader: custom_headers: Additional headers to include in request allow_resume: Whether to support resumable downloads 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: Tuple[bool, str]: (success, save_path or error message) @@ -610,6 +616,12 @@ class Downloader: logger.warning( 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" else: logger.error( @@ -902,6 +914,11 @@ class Downloader: 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: logger.error(f"Unexpected download error: {e}") return False, str(e) @@ -931,6 +948,7 @@ class Downloader: use_auth: bool = False, custom_headers: Optional[Dict[str, str]] = None, return_headers: bool = False, + raise_on_rate_limit: bool = False, ) -> Tuple[bool, Union[bytes, str], Optional[Dict[str, Any]]]: """ Download a file to memory (for small files like preview images) @@ -940,6 +958,10 @@ class Downloader: use_auth: Whether to include authentication headers custom_headers: Additional headers to include in request 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: 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", 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 else: error_msg = f"Download failed with status {response.status}" return False, error_msg, None + except DownloadRateLimitError: + # Structured rate-limit errors must reach the caller unmodified. + raise + except Exception as e: if guard.is_network_unreachable_error(e): guard.register_network_failure(e, destination) diff --git a/py/services/errors.py b/py/services/errors.py index 930febcd..56c26ae4 100644 --- a/py/services/errors.py +++ b/py/services/errors.py @@ -20,6 +20,25 @@ class RateLimitError(RuntimeError): 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): """Raised when a remote resource is permanently missing.""" diff --git a/tests/routes/test_download_queue_handlers.py b/tests/routes/test_download_queue_handlers.py index a838426a..8301b909 100644 --- a/tests/routes/test_download_queue_handlers.py +++ b/tests/routes/test_download_queue_handlers.py @@ -9,6 +9,8 @@ with ``success: false``, not 404. The browser extension's apiFetch treats any import json import logging from pathlib import Path +from types import SimpleNamespace +from unittest.mock import AsyncMock import pytest from aiohttp import web @@ -184,3 +186,83 @@ async def test_retry_failed_history_returns_success( queue = await queue_service.get_queue() assert len(queue) == 1 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 diff --git a/tests/services/test_download_manager_concurrent.py b/tests/services/test_download_manager_concurrent.py index 68a42ef0..fff73887 100644 --- a/tests/services/test_download_manager_concurrent.py +++ b/tests/services/test_download_manager_concurrent.py @@ -130,7 +130,7 @@ async def test_execute_download_uses_rewritten_civitai_preview(monkeypatch, tmp_ self.file_calls: list[tuple[str, str]] = [] 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)) if url.endswith(".jpeg"): Path(path).write_bytes(b"preview") @@ -248,7 +248,7 @@ async def test_execute_download_respects_blur_setting(monkeypatch, tmp_path): def __init__(self): 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)) if url.endswith(".safetensors"): Path(path).write_bytes(b"model") diff --git a/tests/services/test_download_manager_error.py b/tests/services/test_download_manager_error.py index 25872dc0..227f9dc8 100644 --- a/tests/services/test_download_manager_error.py +++ b/tests/services/test_download_manager_error.py @@ -12,9 +12,12 @@ from unittest.mock import AsyncMock import pytest from py.services.download_manager import DownloadManager +from py.services.download_queue_service import DownloadQueueService from py.services.downloader import DownloadStreamControl from py.services import download_manager 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.settings_manager import SettingsManager, get_settings_manager 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 async def test_execute_download_retries_urls(monkeypatch, tmp_path): """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): 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)) if len(self.calls) == 1: return False, "first failed" @@ -373,7 +397,7 @@ async def test_execute_download_adjusts_checkpoint_sub_type(monkeypatch, tmp_pat class DummyDownloader: 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") return True, "ok" @@ -1698,7 +1722,9 @@ async def test_concurrent_downloads_with_same_target_path_do_not_destroy_each_ot downloader_paths = [] 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)) if "2665422" in url: # 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 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" diff --git a/tests/services/test_downloader.py b/tests/services/test_downloader.py index 94ee8e74..d03c2ac9 100644 --- a/tests/services/test_downloader.py +++ b/tests/services/test_downloader.py @@ -5,7 +5,20 @@ from typing import Sequence import pytest +from py.services.connectivity_guard import ConnectivityGuard 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: @@ -297,3 +310,105 @@ async def test_disable_netrc_auth_ignores_netrc_file(tmp_path, monkeypatch): _disable_netrc_auth(session) 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