mirror of
https://github.com/willmiao/ComfyUI-Lora-Manager.git
synced 2026-08-06 14:10:13 -03:00
feat(api): add cursor-based pagination to civitai user-models endpoint
This commit is contained in:
@@ -2590,6 +2590,8 @@ class ModelLibraryHandler:
|
||||
status=400,
|
||||
)
|
||||
|
||||
cursor = request.query.get("cursor")
|
||||
|
||||
metadata_provider = await self._metadata_provider_factory()
|
||||
if not metadata_provider:
|
||||
return web.json_response(
|
||||
@@ -2598,7 +2600,7 @@ class ModelLibraryHandler:
|
||||
)
|
||||
|
||||
try:
|
||||
models = await metadata_provider.get_user_models(username)
|
||||
result = await metadata_provider.get_user_models(username, cursor)
|
||||
except NotImplementedError:
|
||||
return web.json_response(
|
||||
{
|
||||
@@ -2608,14 +2610,35 @@ class ModelLibraryHandler:
|
||||
status=501,
|
||||
)
|
||||
|
||||
if models is None:
|
||||
if result is None:
|
||||
return web.json_response(
|
||||
{"success": False, "error": "Failed to fetch user models"},
|
||||
status=502,
|
||||
)
|
||||
|
||||
if isinstance(result, dict):
|
||||
models = result.get("items")
|
||||
next_cursor = result.get("nextCursor")
|
||||
else:
|
||||
# Defensive: tolerate providers that still return a raw list
|
||||
models = result
|
||||
next_cursor = None
|
||||
|
||||
if not isinstance(models, list):
|
||||
models = []
|
||||
if next_cursor is not None and not isinstance(next_cursor, str):
|
||||
next_cursor = str(next_cursor)
|
||||
|
||||
estimated_total = None
|
||||
if cursor is None:
|
||||
get_count = getattr(metadata_provider, "get_creator_model_count", None)
|
||||
if get_count is not None:
|
||||
try:
|
||||
estimated_total = await get_count(username)
|
||||
except Exception: # best-effort only
|
||||
estimated_total = None
|
||||
if not isinstance(estimated_total, int):
|
||||
estimated_total = None
|
||||
|
||||
lora_scanner = await self._service_registry.get_lora_scanner()
|
||||
checkpoint_scanner = await self._service_registry.get_checkpoint_scanner()
|
||||
@@ -2635,6 +2658,7 @@ class ModelLibraryHandler:
|
||||
versions: list[dict] = []
|
||||
history_service = await self._get_download_history_service()
|
||||
model_ids: list[int] = []
|
||||
model_count = 0
|
||||
for model in models:
|
||||
try:
|
||||
model_ids.append(int(model.get("id")))
|
||||
@@ -2668,6 +2692,8 @@ class ModelLibraryHandler:
|
||||
if model_type not in normalized_allowed_types:
|
||||
continue
|
||||
|
||||
model_count += 1
|
||||
|
||||
scanner = type_scanner_map.get(model_type)
|
||||
if scanner is None:
|
||||
return web.json_response(
|
||||
@@ -2733,7 +2759,15 @@ class ModelLibraryHandler:
|
||||
)
|
||||
|
||||
return web.json_response(
|
||||
{"success": True, "username": username, "versions": versions}
|
||||
{
|
||||
"success": True,
|
||||
"username": username,
|
||||
"versions": versions,
|
||||
"modelCount": model_count,
|
||||
"nextCursor": next_cursor,
|
||||
"hasMore": next_cursor is not None,
|
||||
"estimatedTotal": estimated_total,
|
||||
}
|
||||
)
|
||||
except Exception as exc: # pragma: no cover - defensive logging
|
||||
logger.error("Failed to get Civitai user models: %s", exc, exc_info=True)
|
||||
|
||||
@@ -2,6 +2,7 @@ import asyncio
|
||||
import copy
|
||||
import logging
|
||||
import os
|
||||
import time
|
||||
from collections import OrderedDict
|
||||
from typing import Any, Optional, Dict, Tuple, List, Sequence
|
||||
from .connectivity_guard import (
|
||||
@@ -19,6 +20,12 @@ from ..utils.civitai_utils import resolve_license_payload
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# Best-effort cache for creator model counts, keyed by lowercase username.
|
||||
# Values are (monotonic timestamp, count or None); None results are cached
|
||||
# too so repeated failures don't hammer the API.
|
||||
_CREATOR_COUNT_CACHE_TTL_SECONDS = 600
|
||||
_creator_model_count_cache: Dict[str, Tuple[float, Optional[int]]] = {}
|
||||
|
||||
|
||||
class CivitaiClient:
|
||||
_instance = None
|
||||
@@ -743,17 +750,34 @@ class CivitaiClient:
|
||||
|
||||
return all_versions if all_versions else None
|
||||
|
||||
async def get_user_models(self, username: str) -> Optional[List[Dict]]:
|
||||
"""Fetch all models for a specific Civitai user."""
|
||||
async def get_user_models(
|
||||
self, username: str, cursor: Optional[str] = None
|
||||
) -> Optional[Dict[str, Any]]:
|
||||
"""Fetch one page (up to 100 models) for a specific Civitai user.
|
||||
|
||||
Returns ``{"items": [...], "nextCursor": <str|None>}`` on success,
|
||||
or None on failure. Pass ``cursor`` (from a previous response's
|
||||
``nextCursor``) to fetch subsequent pages.
|
||||
"""
|
||||
if not username:
|
||||
return None
|
||||
|
||||
params: Dict[str, Any] = {
|
||||
"username": username,
|
||||
"nsfw": "true",
|
||||
"limit": 100,
|
||||
"sort": "Newest",
|
||||
"period": "AllTime",
|
||||
}
|
||||
if cursor:
|
||||
params["cursor"] = cursor
|
||||
|
||||
try:
|
||||
success, result = await self._make_request(
|
||||
"GET",
|
||||
f"{self.base_url}/models",
|
||||
use_auth=True,
|
||||
params={"username": username, "nsfw": "true"},
|
||||
params=params,
|
||||
)
|
||||
|
||||
if not success:
|
||||
@@ -765,7 +789,7 @@ class CivitaiClient:
|
||||
|
||||
items = result.get("items") if isinstance(result, dict) else None
|
||||
if not isinstance(items, list):
|
||||
return []
|
||||
items = []
|
||||
|
||||
for model in items:
|
||||
versions = model.get("modelVersions")
|
||||
@@ -774,9 +798,68 @@ class CivitaiClient:
|
||||
for version in versions:
|
||||
self._remove_comfy_metadata(version)
|
||||
|
||||
return items
|
||||
next_cursor: Optional[str] = None
|
||||
metadata = result.get("metadata") if isinstance(result, dict) else None
|
||||
if isinstance(metadata, dict):
|
||||
raw_cursor = metadata.get("nextCursor")
|
||||
if raw_cursor is not None:
|
||||
next_cursor = str(raw_cursor)
|
||||
|
||||
return {"items": items, "nextCursor": next_cursor}
|
||||
except RateLimitError:
|
||||
raise
|
||||
except Exception as exc: # pragma: no cover - defensive logging
|
||||
logger.error("Error fetching models for %s: %s", username, exc)
|
||||
return None
|
||||
|
||||
async def get_creator_model_count(self, username: str) -> Optional[int]:
|
||||
"""Best-effort lookup of a creator's published model count.
|
||||
|
||||
Uses the ``/creators`` endpoint (a contains-match query), picking the
|
||||
entry whose username matches exactly (case-insensitive). Returns None
|
||||
on any failure; never raises. Results (including None) are cached
|
||||
for ``_CREATOR_COUNT_CACHE_TTL_SECONDS``.
|
||||
"""
|
||||
if not username:
|
||||
return None
|
||||
|
||||
cache_key = username.lower()
|
||||
cached = _creator_model_count_cache.get(cache_key)
|
||||
if cached is not None:
|
||||
cached_at, cached_count = cached
|
||||
if time.monotonic() - cached_at < _CREATOR_COUNT_CACHE_TTL_SECONDS:
|
||||
return cached_count
|
||||
|
||||
count: Optional[int] = None
|
||||
try:
|
||||
success, result = await self._make_request(
|
||||
"GET",
|
||||
f"{self.base_url}/creators",
|
||||
use_auth=True,
|
||||
params={"query": username, "limit": 10},
|
||||
)
|
||||
|
||||
if success and isinstance(result, dict):
|
||||
creators = result.get("items")
|
||||
if isinstance(creators, list):
|
||||
for creator in creators:
|
||||
if not isinstance(creator, dict):
|
||||
continue
|
||||
creator_name = creator.get("username")
|
||||
if not isinstance(creator_name, str):
|
||||
continue
|
||||
if creator_name.lower() != cache_key:
|
||||
continue
|
||||
model_count = creator.get("modelCount")
|
||||
if isinstance(model_count, (int, float)) and not isinstance(
|
||||
model_count, bool
|
||||
):
|
||||
count = int(model_count)
|
||||
break
|
||||
except Exception as exc: # best-effort only, never propagate
|
||||
logger.debug(
|
||||
"Failed to fetch creator model count for %s: %s", username, exc
|
||||
)
|
||||
|
||||
_creator_model_count_cache[cache_key] = (time.monotonic(), count)
|
||||
return count
|
||||
|
||||
@@ -143,10 +143,18 @@ class ModelMetadataProvider(ABC):
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
async def get_user_models(self, username: str) -> Optional[List[Dict]]:
|
||||
"""Fetch models owned by the specified user"""
|
||||
async def get_user_models(self, username: str, cursor: Optional[str] = None) -> Optional[Dict]:
|
||||
"""Fetch one page of models owned by the specified user.
|
||||
|
||||
Returns ``{"items": [...], "nextCursor": <str|None>}`` on success,
|
||||
or None when unsupported/failed. ``cursor`` continues a previous page.
|
||||
"""
|
||||
pass
|
||||
|
||||
async def get_creator_model_count(self, username: str) -> Optional[int]:
|
||||
"""Published model count for the user; None when unsupported."""
|
||||
return None
|
||||
|
||||
class CivitaiModelMetadataProvider(ModelMetadataProvider):
|
||||
"""Provider that uses Civitai API for metadata"""
|
||||
|
||||
@@ -175,8 +183,11 @@ class CivitaiModelMetadataProvider(ModelMetadataProvider):
|
||||
async def get_model_version_info(self, version_id: str) -> Tuple[Optional[Dict], Optional[str]]:
|
||||
return await self.client.get_model_version_info(version_id)
|
||||
|
||||
async def get_user_models(self, username: str) -> Optional[List[Dict]]:
|
||||
return await self.client.get_user_models(username)
|
||||
async def get_user_models(self, username: str, cursor: Optional[str] = None) -> Optional[Dict]:
|
||||
return await self.client.get_user_models(username, cursor)
|
||||
|
||||
async def get_creator_model_count(self, username: str) -> Optional[int]:
|
||||
return await self.client.get_creator_model_count(username)
|
||||
|
||||
class CivArchiveModelMetadataProvider(ModelMetadataProvider):
|
||||
"""Provider that uses CivArchive API for metadata"""
|
||||
@@ -196,7 +207,7 @@ class CivArchiveModelMetadataProvider(ModelMetadataProvider):
|
||||
async def get_model_version_info(self, version_id: str) -> Tuple[Optional[Dict], Optional[str]]:
|
||||
return await self.client.get_model_version_info(version_id)
|
||||
|
||||
async def get_user_models(self, username: str) -> Optional[List[Dict]]:
|
||||
async def get_user_models(self, username: str, cursor: Optional[str] = None) -> Optional[Dict]:
|
||||
"""Not supported by CivArchive provider"""
|
||||
return None
|
||||
|
||||
@@ -347,7 +358,7 @@ class SQLiteModelMetadataProvider(ModelMetadataProvider):
|
||||
version_data = await self._get_version_with_model_data(db, model_id, version_id)
|
||||
return version_data, None
|
||||
|
||||
async def get_user_models(self, username: str) -> Optional[List[Dict]]:
|
||||
async def get_user_models(self, username: str, cursor: Optional[str] = None) -> Optional[Dict]:
|
||||
"""Listing models by username is not supported for archive database"""
|
||||
return None
|
||||
|
||||
@@ -602,13 +613,14 @@ class FallbackMetadataProvider(ModelMetadataProvider):
|
||||
continue
|
||||
return None
|
||||
|
||||
async def get_user_models(self, username: str) -> Optional[List[Dict]]:
|
||||
async def get_user_models(self, username: str, cursor: Optional[str] = None) -> Optional[Dict]:
|
||||
for provider, label in self._iter_providers():
|
||||
try:
|
||||
result = await self._call_with_rate_limit(
|
||||
label,
|
||||
provider.get_user_models,
|
||||
username,
|
||||
cursor=cursor,
|
||||
)
|
||||
if result is not None:
|
||||
return result
|
||||
@@ -624,6 +636,19 @@ class FallbackMetadataProvider(ModelMetadataProvider):
|
||||
continue
|
||||
return None
|
||||
|
||||
async def get_creator_model_count(self, username: str) -> Optional[int]:
|
||||
for provider, label in self._iter_providers():
|
||||
try:
|
||||
result = await provider.get_creator_model_count(username)
|
||||
if result is not None:
|
||||
return result
|
||||
except Exception as e:
|
||||
logger.debug(
|
||||
"Provider %s failed for get_creator_model_count: %s", label, e
|
||||
)
|
||||
continue
|
||||
return None
|
||||
|
||||
def _iter_providers(self):
|
||||
return zip(self.providers, self._provider_labels)
|
||||
|
||||
@@ -704,13 +729,17 @@ class RateLimitRetryingProvider(ModelMetadataProvider):
|
||||
version_id,
|
||||
)
|
||||
|
||||
async def get_user_models(self, username: str) -> Optional[List[Dict]]:
|
||||
async def get_user_models(self, username: str, cursor: Optional[str] = None) -> Optional[Dict]:
|
||||
return await self._rate_limit_helper.run(
|
||||
self._label,
|
||||
self._provider.get_user_models,
|
||||
username,
|
||||
cursor=cursor,
|
||||
)
|
||||
|
||||
async def get_creator_model_count(self, username: str) -> Optional[int]:
|
||||
return await self._provider.get_creator_model_count(username)
|
||||
|
||||
class ModelMetadataProviderManager:
|
||||
"""Manager for selecting and using model metadata providers"""
|
||||
|
||||
@@ -776,10 +805,20 @@ class ModelMetadataProviderManager:
|
||||
except NotImplementedError:
|
||||
return None
|
||||
|
||||
async def get_user_models(self, username: str, provider_name: str = None) -> Optional[List[Dict]]:
|
||||
"""Fetch models owned by the specified user"""
|
||||
async def get_user_models(
|
||||
self,
|
||||
username: str,
|
||||
provider_name: str = None,
|
||||
cursor: Optional[str] = None,
|
||||
) -> Optional[Dict]:
|
||||
"""Fetch one page of models owned by the specified user"""
|
||||
provider = self._get_provider(provider_name)
|
||||
return await provider.get_user_models(username)
|
||||
return await provider.get_user_models(username, cursor)
|
||||
|
||||
async def get_creator_model_count(self, username: str, provider_name: str = None) -> Optional[int]:
|
||||
"""Best-effort published model count for the specified user"""
|
||||
provider = self._get_provider(provider_name)
|
||||
return await provider.get_creator_model_count(username)
|
||||
|
||||
def _get_provider(self, provider_name: str = None) -> ModelMetadataProvider:
|
||||
"""Get provider by name or default provider"""
|
||||
|
||||
@@ -900,18 +900,28 @@ class FakeMetadataProvider:
|
||||
async def get_model_versions(self, _model_id):
|
||||
return {"modelVersions": [], "name": "", "type": "lora"}
|
||||
|
||||
async def get_user_models(self, _username):
|
||||
return []
|
||||
async def get_user_models(self, _username, cursor=None):
|
||||
return {"items": [], "nextCursor": None}
|
||||
|
||||
async def get_creator_model_count(self, _username):
|
||||
return None
|
||||
|
||||
|
||||
class FakeUserModelsProvider(FakeMetadataProvider):
|
||||
def __init__(self, models):
|
||||
def __init__(self, models, next_cursor=None, estimated_total=None):
|
||||
self.models = models
|
||||
self.next_cursor = next_cursor
|
||||
self.estimated_total = estimated_total
|
||||
self.received_usernames: list[str] = []
|
||||
self.received_cursors: list = []
|
||||
|
||||
async def get_user_models(self, username):
|
||||
async def get_user_models(self, username, cursor=None):
|
||||
self.received_usernames.append(username)
|
||||
return self.models
|
||||
self.received_cursors.append(cursor)
|
||||
return {"items": self.models, "nextCursor": self.next_cursor}
|
||||
|
||||
async def get_creator_model_count(self, _username):
|
||||
return self.estimated_total
|
||||
|
||||
|
||||
async def fake_metadata_provider_factory():
|
||||
@@ -1286,6 +1296,88 @@ async def test_get_civitai_user_models_requires_username():
|
||||
assert "username" in payload["error"].lower()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_civitai_user_models_returns_pagination_fields():
|
||||
models = [
|
||||
{
|
||||
"id": 1,
|
||||
"name": "Model A",
|
||||
"type": "LORA",
|
||||
"tags": [],
|
||||
"modelVersions": [
|
||||
{"id": 100, "name": "v1", "images": [{"url": "http://example.com/a.jpg"}]},
|
||||
],
|
||||
},
|
||||
{
|
||||
"id": 2,
|
||||
"name": "Unsupported",
|
||||
"type": "Other",
|
||||
"modelVersions": [{"id": 200, "name": "v1"}],
|
||||
},
|
||||
]
|
||||
|
||||
provider = FakeUserModelsProvider(models, next_cursor="cursor-token", estimated_total=2140)
|
||||
|
||||
async def provider_factory():
|
||||
return provider
|
||||
|
||||
handler = ModelLibraryHandler(
|
||||
ServiceRegistryAdapter(
|
||||
get_lora_scanner=fake_scanner_factory,
|
||||
get_checkpoint_scanner=fake_scanner_factory,
|
||||
get_embedding_scanner=fake_scanner_factory,
|
||||
get_downloaded_version_history_service=fake_download_history_service_factory,
|
||||
),
|
||||
metadata_provider_factory=provider_factory,
|
||||
)
|
||||
|
||||
response = await handler.get_civitai_user_models(
|
||||
FakeRequest(query={"username": "pixel"})
|
||||
)
|
||||
payload = json.loads(response.text)
|
||||
|
||||
assert response.status == 200
|
||||
assert payload["success"] is True
|
||||
# modelCount only counts models surviving the type filter
|
||||
assert payload["modelCount"] == 1
|
||||
assert payload["nextCursor"] == "cursor-token"
|
||||
assert payload["hasMore"] is True
|
||||
# first page includes the estimated total
|
||||
assert payload["estimatedTotal"] == 2140
|
||||
assert provider.received_cursors == [None]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_civitai_user_models_passes_cursor_and_omits_estimate():
|
||||
provider = FakeUserModelsProvider([], next_cursor=None, estimated_total=999)
|
||||
|
||||
async def provider_factory():
|
||||
return provider
|
||||
|
||||
handler = ModelLibraryHandler(
|
||||
ServiceRegistryAdapter(
|
||||
get_lora_scanner=fake_scanner_factory,
|
||||
get_checkpoint_scanner=fake_scanner_factory,
|
||||
get_embedding_scanner=fake_scanner_factory,
|
||||
get_downloaded_version_history_service=fake_download_history_service_factory,
|
||||
),
|
||||
metadata_provider_factory=provider_factory,
|
||||
)
|
||||
|
||||
response = await handler.get_civitai_user_models(
|
||||
FakeRequest(query={"username": "pixel", "cursor": "opaque-token"})
|
||||
)
|
||||
payload = json.loads(response.text)
|
||||
|
||||
assert response.status == 200
|
||||
assert payload["success"] is True
|
||||
assert payload["nextCursor"] is None
|
||||
assert payload["hasMore"] is False
|
||||
# cursor requests must not include the estimated total
|
||||
assert payload["estimatedTotal"] is None
|
||||
assert provider.received_cursors == ["opaque-token"]
|
||||
|
||||
|
||||
def test_ensure_handler_mapping_caches_result():
|
||||
call_records = []
|
||||
|
||||
|
||||
@@ -35,9 +35,11 @@ class DummyDownloader:
|
||||
def reset_singletons():
|
||||
CivitaiClient._instance = None
|
||||
ModelMetadataProviderManager._instance = None
|
||||
civitai_client_module._creator_model_count_cache.clear()
|
||||
yield
|
||||
CivitaiClient._instance = None
|
||||
ModelMetadataProviderManager._instance = None
|
||||
civitai_client_module._creator_model_count_cache.clear()
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
@@ -622,3 +624,162 @@ async def test_get_image_info_handles_invalid_id(monkeypatch, downloader, caplog
|
||||
|
||||
assert result is None
|
||||
assert "Invalid image ID format" in caplog.text
|
||||
|
||||
|
||||
async def test_get_user_models_requests_first_page_with_stable_params(downloader):
|
||||
request_calls = []
|
||||
|
||||
async def fake_make_request(method, url, use_auth=True, **kwargs):
|
||||
request_calls.append({"method": method, "url": url, "kwargs": kwargs})
|
||||
return True, {
|
||||
"items": [
|
||||
{
|
||||
"id": 1,
|
||||
"modelVersions": [
|
||||
{"id": 100, "images": [{"meta": {"comfy": {"x": 1}}}]}
|
||||
],
|
||||
}
|
||||
],
|
||||
"metadata": {"nextCursor": "next-token"},
|
||||
}
|
||||
|
||||
downloader.make_request = fake_make_request
|
||||
|
||||
client = await CivitaiClient.get_instance()
|
||||
result = await client.get_user_models("pixel")
|
||||
|
||||
assert result is not None
|
||||
assert result["nextCursor"] == "next-token"
|
||||
assert len(result["items"]) == 1
|
||||
# comfy metadata is still stripped
|
||||
assert "comfy" not in result["items"][0]["modelVersions"][0]["images"][0]["meta"]
|
||||
|
||||
call = request_calls[0]
|
||||
assert call["method"] == "GET"
|
||||
assert call["url"] == "https://civitai.red/api/v1/models"
|
||||
params = call["kwargs"]["params"]
|
||||
assert params["username"] == "pixel"
|
||||
assert params["nsfw"] == "true"
|
||||
assert params["limit"] == 100
|
||||
assert params["sort"] == "Newest"
|
||||
assert params["period"] == "AllTime"
|
||||
assert "cursor" not in params
|
||||
|
||||
|
||||
async def test_get_user_models_passes_cursor_and_stringifies_next_cursor(downloader):
|
||||
request_calls = []
|
||||
|
||||
async def fake_make_request(method, url, use_auth=True, **kwargs):
|
||||
request_calls.append(kwargs)
|
||||
return True, {"items": [], "metadata": {"nextCursor": 12345}}
|
||||
|
||||
downloader.make_request = fake_make_request
|
||||
|
||||
client = await CivitaiClient.get_instance()
|
||||
result = await client.get_user_models("pixel", cursor="opaque-token")
|
||||
|
||||
assert request_calls[0]["params"]["cursor"] == "opaque-token"
|
||||
assert result == {"items": [], "nextCursor": "12345"}
|
||||
|
||||
|
||||
async def test_get_user_models_without_next_cursor_returns_none_cursor(downloader):
|
||||
async def fake_make_request(method, url, use_auth=True, **kwargs):
|
||||
return True, {"items": [{"id": 1, "modelVersions": []}], "metadata": {}}
|
||||
|
||||
downloader.make_request = fake_make_request
|
||||
|
||||
client = await CivitaiClient.get_instance()
|
||||
result = await client.get_user_models("pixel")
|
||||
|
||||
assert result == {"items": [{"id": 1, "modelVersions": []}], "nextCursor": None}
|
||||
|
||||
|
||||
async def test_get_user_models_failure_returns_none(downloader):
|
||||
async def fake_make_request(method, url, use_auth=True, **kwargs):
|
||||
return False, "500 server error"
|
||||
|
||||
downloader.make_request = fake_make_request
|
||||
|
||||
client = await CivitaiClient.get_instance()
|
||||
result = await client.get_user_models("pixel")
|
||||
|
||||
assert result is None
|
||||
|
||||
|
||||
async def test_get_creator_model_count_matches_exact_username(downloader):
|
||||
request_calls = []
|
||||
|
||||
async def fake_make_request(method, url, use_auth=True, **kwargs):
|
||||
request_calls.append({"url": url, "kwargs": kwargs})
|
||||
return True, {
|
||||
"items": [
|
||||
{"username": "pixelart", "modelCount": 5},
|
||||
{"username": "Pixel", "modelCount": 2140},
|
||||
]
|
||||
}
|
||||
|
||||
downloader.make_request = fake_make_request
|
||||
|
||||
client = await CivitaiClient.get_instance()
|
||||
count = await client.get_creator_model_count("pixel")
|
||||
|
||||
assert count == 2140
|
||||
assert request_calls[0]["url"] == "https://civitai.red/api/v1/creators"
|
||||
assert request_calls[0]["kwargs"]["params"] == {"query": "pixel", "limit": 10}
|
||||
|
||||
|
||||
async def test_get_creator_model_count_without_exact_match_returns_none(downloader):
|
||||
async def fake_make_request(method, url, use_auth=True, **kwargs):
|
||||
return True, {"items": [{"username": "pixelart", "modelCount": 5}]}
|
||||
|
||||
downloader.make_request = fake_make_request
|
||||
|
||||
client = await CivitaiClient.get_instance()
|
||||
count = await client.get_creator_model_count("pixel")
|
||||
|
||||
assert count is None
|
||||
|
||||
|
||||
async def test_get_creator_model_count_caches_results(downloader):
|
||||
request_count = 0
|
||||
|
||||
async def fake_make_request(method, url, use_auth=True, **kwargs):
|
||||
nonlocal request_count
|
||||
request_count += 1
|
||||
return True, {"items": [{"username": "pixel", "modelCount": 42}]}
|
||||
|
||||
downloader.make_request = fake_make_request
|
||||
|
||||
client = await CivitaiClient.get_instance()
|
||||
|
||||
assert await client.get_creator_model_count("pixel") == 42
|
||||
# case-insensitive cache key, second call served from cache
|
||||
assert await client.get_creator_model_count("Pixel") == 42
|
||||
assert request_count == 1
|
||||
|
||||
|
||||
async def test_get_creator_model_count_caches_failures(downloader):
|
||||
request_count = 0
|
||||
|
||||
async def fake_make_request(method, url, use_auth=True, **kwargs):
|
||||
nonlocal request_count
|
||||
request_count += 1
|
||||
return False, "500 server error"
|
||||
|
||||
downloader.make_request = fake_make_request
|
||||
|
||||
client = await CivitaiClient.get_instance()
|
||||
|
||||
assert await client.get_creator_model_count("pixel") is None
|
||||
assert await client.get_creator_model_count("pixel") is None
|
||||
assert request_count == 1
|
||||
|
||||
|
||||
async def test_get_creator_model_count_never_raises(downloader):
|
||||
async def fake_make_request(method, url, use_auth=True, **kwargs):
|
||||
return True, "unexpected non-dict payload"
|
||||
|
||||
downloader.make_request = fake_make_request
|
||||
|
||||
client = await CivitaiClient.get_instance()
|
||||
assert await client.get_creator_model_count("pixel") is None
|
||||
|
||||
Reference in New Issue
Block a user