mirror of
https://github.com/willmiao/ComfyUI-Lora-Manager.git
synced 2026-10-03 00:35:32 -03:00
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:
@@ -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
|
||||
|
||||
@@ -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")
|
||||
|
||||
@@ -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"
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user