diff --git a/py/services/download_manager.py b/py/services/download_manager.py index 29dcca9b..dc3d1a5a 100644 --- a/py/services/download_manager.py +++ b/py/services/download_manager.py @@ -70,8 +70,11 @@ CIVITAI_DOWNLOAD_URL_PREFIXES = ( # 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") +# occupies a concurrency slot. civarchive.com is only gated for downloads +# whose source is CivArchive: a cooldown armed by background metadata fetches +# must not block plain CivitAI downloads. +DOWNLOAD_PREFLIGHT_HOSTS = ("civitai.com", "civitai.red") +DOWNLOAD_PREFLIGHT_HOSTS_CIVARCHIVE = DOWNLOAD_PREFLIGHT_HOSTS + ("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). @@ -258,19 +261,30 @@ class DownloadManager: hostname = urlparse(url).hostname return hostname.lower() if hostname else "unknown" - async def _preflight_rate_limit_error(self) -> Optional[DownloadRateLimitError]: + async def _preflight_rate_limit_error( + self, source: str | None = None + ) -> 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 + hosts this 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). + + civarchive.com is only consulted when ``source == "civarchive"`` — + other downloads never touch it, so a cooldown armed there (typically + by bulk metadata fetches) must not block them. """ + hosts = ( + DOWNLOAD_PREFLIGHT_HOSTS_CIVARCHIVE + if source == "civarchive" + else DOWNLOAD_PREFLIGHT_HOSTS + ) coordinator = await RateLimitCoordinator.get_instance() worst_host: Optional[str] = None worst_remaining = 0.0 - for host in DOWNLOAD_PREFLIGHT_HOSTS: + for host in hosts: remaining = coordinator.remaining_seconds(host) if remaining > worst_remaining: worst_host = host @@ -678,7 +692,7 @@ class DownloadManager: # 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() + preflight_error = await self._preflight_rate_limit_error(source) if preflight_error is not None: logger.info( "Download %s skipped: %s", task_id, preflight_error diff --git a/tests/services/test_download_manager_error.py b/tests/services/test_download_manager_error.py index 227f9dc8..6ace864b 100644 --- a/tests/services/test_download_manager_error.py +++ b/tests/services/test_download_manager_error.py @@ -1969,7 +1969,7 @@ def _prepare_tracked_download(manager, download_id, status="downloading"): manager._pause_events[download_id] = DownloadStreamControl() -async def _run_download(manager, download_id, tmp_path): +async def _run_download(manager, download_id, tmp_path, source=None): return await manager._download_with_semaphore( download_id, 1, @@ -1978,7 +1978,7 @@ async def _run_download(manager, download_id, tmp_path): "", None, False, - None, + source, None, False, ) @@ -2063,6 +2063,59 @@ async def test_preflight_gate_blocks_when_host_in_cooldown( manager._download_semaphore.release() +@pytest.mark.asyncio +async def test_preflight_gate_ignores_civarchive_cooldown_for_civitai_downloads( + monkeypatch, tmp_path, queue_service, reset_rate_limit_coordinator +): + """A civarchive.com cooldown (armed e.g. by bulk metadata fetches) must + not block a plain CivitAI download that never touches that host.""" + manager = DownloadManager() + monkeypatch.setattr(manager, "_cleanup_download_record", AsyncMock()) + + coordinator = await RateLimitCoordinator.get_instance() + coordinator.register_rate_limit("civarchive.com", 1800) + + download_id = "dl-civarchive-cooldown" + await queue_service.add_to_queue(download_id=download_id, model_id=1) + _prepare_tracked_download(manager, download_id, status="waiting") + + execute = AsyncMock(return_value={"success": True}) + monkeypatch.setattr(manager, "_execute_original_download", execute) + + result = await _run_download(manager, download_id, tmp_path) + + assert execute.await_count == 1 + assert result["success"] is True + + +@pytest.mark.asyncio +async def test_preflight_gate_blocks_civarchive_source_during_civarchive_cooldown( + monkeypatch, tmp_path, queue_service, reset_rate_limit_coordinator +): + """Downloads explicitly sourced from CivArchive still gate on its cooldown.""" + manager = DownloadManager() + monkeypatch.setattr(manager, "_cleanup_download_record", AsyncMock()) + + coordinator = await RateLimitCoordinator.get_instance() + coordinator.register_rate_limit("civarchive.com", 1800) + + download_id = "dl-civarchive-source" + 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, source="civarchive") + + assert execute.await_count == 0 + assert result["success"] is False + assert result["reason"] == "rate_limited" + assert "civarchive.com" in result["error"] + + @pytest.mark.asyncio async def test_rate_limit_retry_after_falls_back_to_coordinator_backoff( monkeypatch, tmp_path, queue_service, reset_rate_limit_coordinator