mirror of
https://github.com/willmiao/ComfyUI-Lora-Manager.git
synced 2026-08-06 22:10:14 -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,
|
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)
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -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"""
|
||||||
|
|||||||
@@ -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 = []
|
||||||
|
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
Reference in New Issue
Block a user