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,
)
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)

View File

@@ -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

View File

@@ -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"""

View File

@@ -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 = []

View File

@@ -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