feat(api): add cursor-based pagination to civitai user-models endpoint

This commit is contained in:
Will Miao
2026-08-03 11:07:06 +08:00
parent 823f71f269
commit 75e63c758b
5 changed files with 433 additions and 24 deletions

View File

@@ -2590,6 +2590,8 @@ class ModelLibraryHandler:
status=400, status=400,
) )
cursor = request.query.get("cursor")
metadata_provider = await self._metadata_provider_factory() metadata_provider = await self._metadata_provider_factory()
if not metadata_provider: if not metadata_provider:
return web.json_response( return web.json_response(
@@ -2598,7 +2600,7 @@ class ModelLibraryHandler:
) )
try: try:
models = await metadata_provider.get_user_models(username) result = await metadata_provider.get_user_models(username, cursor)
except NotImplementedError: except NotImplementedError:
return web.json_response( return web.json_response(
{ {
@@ -2608,14 +2610,35 @@ class ModelLibraryHandler:
status=501, status=501,
) )
if models is None: if result is None:
return web.json_response( return web.json_response(
{"success": False, "error": "Failed to fetch user models"}, {"success": False, "error": "Failed to fetch user models"},
status=502, 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): if not isinstance(models, list):
models = [] 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() lora_scanner = await self._service_registry.get_lora_scanner()
checkpoint_scanner = await self._service_registry.get_checkpoint_scanner() checkpoint_scanner = await self._service_registry.get_checkpoint_scanner()
@@ -2635,6 +2658,7 @@ class ModelLibraryHandler:
versions: list[dict] = [] versions: list[dict] = []
history_service = await self._get_download_history_service() history_service = await self._get_download_history_service()
model_ids: list[int] = [] model_ids: list[int] = []
model_count = 0
for model in models: for model in models:
try: try:
model_ids.append(int(model.get("id"))) model_ids.append(int(model.get("id")))
@@ -2668,6 +2692,8 @@ class ModelLibraryHandler:
if model_type not in normalized_allowed_types: if model_type not in normalized_allowed_types:
continue continue
model_count += 1
scanner = type_scanner_map.get(model_type) scanner = type_scanner_map.get(model_type)
if scanner is None: if scanner is None:
return web.json_response( return web.json_response(
@@ -2733,7 +2759,15 @@ class ModelLibraryHandler:
) )
return web.json_response( 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 except Exception as exc: # pragma: no cover - defensive logging
logger.error("Failed to get Civitai user models: %s", exc, exc_info=True) logger.error("Failed to get Civitai user models: %s", exc, exc_info=True)

View File

@@ -2,6 +2,7 @@ import asyncio
import copy import copy
import logging import logging
import os import os
import time
from collections import OrderedDict from collections import OrderedDict
from typing import Any, Optional, Dict, Tuple, List, Sequence from typing import Any, Optional, Dict, Tuple, List, Sequence
from .connectivity_guard import ( from .connectivity_guard import (
@@ -19,6 +20,12 @@ from ..utils.civitai_utils import resolve_license_payload
logger = logging.getLogger(__name__) 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: class CivitaiClient:
_instance = None _instance = None
@@ -743,17 +750,34 @@ class CivitaiClient:
return all_versions if all_versions else None return all_versions if all_versions else None
async def get_user_models(self, username: str) -> Optional[List[Dict]]: async def get_user_models(
"""Fetch all models for a specific Civitai user.""" 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: if not username:
return None return None
params: Dict[str, Any] = {
"username": username,
"nsfw": "true",
"limit": 100,
"sort": "Newest",
"period": "AllTime",
}
if cursor:
params["cursor"] = cursor
try: try:
success, result = await self._make_request( success, result = await self._make_request(
"GET", "GET",
f"{self.base_url}/models", f"{self.base_url}/models",
use_auth=True, use_auth=True,
params={"username": username, "nsfw": "true"}, params=params,
) )
if not success: if not success:
@@ -765,7 +789,7 @@ class CivitaiClient:
items = result.get("items") if isinstance(result, dict) else None items = result.get("items") if isinstance(result, dict) else None
if not isinstance(items, list): if not isinstance(items, list):
return [] items = []
for model in items: for model in items:
versions = model.get("modelVersions") versions = model.get("modelVersions")
@@ -774,9 +798,68 @@ class CivitaiClient:
for version in versions: for version in versions:
self._remove_comfy_metadata(version) 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: except RateLimitError:
raise raise
except Exception as exc: # pragma: no cover - defensive logging except Exception as exc: # pragma: no cover - defensive logging
logger.error("Error fetching models for %s: %s", username, exc) logger.error("Error fetching models for %s: %s", username, exc)
return None 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

View File

@@ -143,10 +143,18 @@ class ModelMetadataProvider(ABC):
pass pass
@abstractmethod @abstractmethod
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]:
"""Fetch models owned by the specified user""" """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 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): class CivitaiModelMetadataProvider(ModelMetadataProvider):
"""Provider that uses Civitai API for metadata""" """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]]: 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) 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]:
return await self.client.get_user_models(username) 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): class CivArchiveModelMetadataProvider(ModelMetadataProvider):
"""Provider that uses CivArchive API for metadata""" """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]]: 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) 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""" """Not supported by CivArchive provider"""
return None return None
@@ -347,7 +358,7 @@ class SQLiteModelMetadataProvider(ModelMetadataProvider):
version_data = await self._get_version_with_model_data(db, model_id, version_id) version_data = await self._get_version_with_model_data(db, model_id, version_id)
return version_data, None 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""" """Listing models by username is not supported for archive database"""
return None return None
@@ -602,13 +613,14 @@ class FallbackMetadataProvider(ModelMetadataProvider):
continue continue
return None 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(): for provider, label in self._iter_providers():
try: try:
result = await self._call_with_rate_limit( result = await self._call_with_rate_limit(
label, label,
provider.get_user_models, provider.get_user_models,
username, username,
cursor=cursor,
) )
if result is not None: if result is not None:
return result return result
@@ -624,6 +636,19 @@ class FallbackMetadataProvider(ModelMetadataProvider):
continue continue
return None 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): def _iter_providers(self):
return zip(self.providers, self._provider_labels) return zip(self.providers, self._provider_labels)
@@ -704,13 +729,17 @@ class RateLimitRetryingProvider(ModelMetadataProvider):
version_id, 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( return await self._rate_limit_helper.run(
self._label, self._label,
self._provider.get_user_models, self._provider.get_user_models,
username, 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: class ModelMetadataProviderManager:
"""Manager for selecting and using model metadata providers""" """Manager for selecting and using model metadata providers"""
@@ -776,10 +805,20 @@ class ModelMetadataProviderManager:
except NotImplementedError: except NotImplementedError:
return None return None
async def get_user_models(self, username: str, provider_name: str = None) -> Optional[List[Dict]]: async def get_user_models(
"""Fetch models owned by the specified user""" 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) 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: def _get_provider(self, provider_name: str = None) -> ModelMetadataProvider:
"""Get provider by name or default provider""" """Get provider by name or default provider"""

View File

@@ -900,18 +900,28 @@ class FakeMetadataProvider:
async def get_model_versions(self, _model_id): async def get_model_versions(self, _model_id):
return {"modelVersions": [], "name": "", "type": "lora"} return {"modelVersions": [], "name": "", "type": "lora"}
async def get_user_models(self, _username): async def get_user_models(self, _username, cursor=None):
return [] return {"items": [], "nextCursor": None}
async def get_creator_model_count(self, _username):
return None
class FakeUserModelsProvider(FakeMetadataProvider): class FakeUserModelsProvider(FakeMetadataProvider):
def __init__(self, models): def __init__(self, models, next_cursor=None, estimated_total=None):
self.models = models self.models = models
self.next_cursor = next_cursor
self.estimated_total = estimated_total
self.received_usernames: list[str] = [] 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) 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(): 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() 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(): def test_ensure_handler_mapping_caches_result():
call_records = [] call_records = []

View File

@@ -35,9 +35,11 @@ class DummyDownloader:
def reset_singletons(): def reset_singletons():
CivitaiClient._instance = None CivitaiClient._instance = None
ModelMetadataProviderManager._instance = None ModelMetadataProviderManager._instance = None
civitai_client_module._creator_model_count_cache.clear()
yield yield
CivitaiClient._instance = None CivitaiClient._instance = None
ModelMetadataProviderManager._instance = None ModelMetadataProviderManager._instance = None
civitai_client_module._creator_model_count_cache.clear()
@pytest.fixture @pytest.fixture
@@ -622,3 +624,162 @@ async def test_get_image_info_handles_invalid_id(monkeypatch, downloader, caplog
assert result is None assert result is None
assert "Invalid image ID format" in caplog.text 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