From 75e63c758bf0dea326fad6e6772592f5b7e858d9 Mon Sep 17 00:00:00 2001 From: Will Miao Date: Mon, 3 Aug 2026 11:07:06 +0800 Subject: [PATCH] feat(api): add cursor-based pagination to civitai user-models endpoint --- py/routes/handlers/misc_handlers.py | 40 +++++- py/services/civitai_client.py | 93 +++++++++++++- py/services/model_metadata_provider.py | 61 ++++++++-- tests/routes/test_misc_routes.py | 102 +++++++++++++++- tests/services/test_civitai_client.py | 161 +++++++++++++++++++++++++ 5 files changed, 433 insertions(+), 24 deletions(-) diff --git a/py/routes/handlers/misc_handlers.py b/py/routes/handlers/misc_handlers.py index 842d48b3..6cc2da27 100644 --- a/py/routes/handlers/misc_handlers.py +++ b/py/routes/handlers/misc_handlers.py @@ -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) diff --git a/py/services/civitai_client.py b/py/services/civitai_client.py index 788927c7..27f581ee 100644 --- a/py/services/civitai_client.py +++ b/py/services/civitai_client.py @@ -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": }`` 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 diff --git a/py/services/model_metadata_provider.py b/py/services/model_metadata_provider.py index b22300df..999c7936 100644 --- a/py/services/model_metadata_provider.py +++ b/py/services/model_metadata_provider.py @@ -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": }`` 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""" diff --git a/tests/routes/test_misc_routes.py b/tests/routes/test_misc_routes.py index b6a7566d..a735a4ff 100644 --- a/tests/routes/test_misc_routes.py +++ b/tests/routes/test_misc_routes.py @@ -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 = [] diff --git a/tests/services/test_civitai_client.py b/tests/services/test_civitai_client.py index 0203b457..cd56d15d 100644 --- a/tests/services/test_civitai_client.py +++ b/tests/services/test_civitai_client.py @@ -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