merge: modules/civitai

This commit is contained in:
vladmandic
2026-03-17 11:07:51 +01:00
parent 4c256976df
commit ea1abfe2ce
9 changed files with 1707 additions and 459 deletions
+589 -47
View File
@@ -1,68 +1,610 @@
"""CivitAI REST API endpoints.
Provides search, model/version lookup, download queue management,
options discovery, metadata scanning, and user data (bookmarks/bans/history).
"""
import os
from starlette.responses import JSONResponse
from modules.logger import log
def models_to_json(all_models:list, model_id:int=None):
dct = []
for model in all_models:
if model_id is not None and model.id != model_id:
continue
model_dct = model.__dict__.copy()
versions_dct = []
for version in model.versions:
version_dct = version.__dict__.copy()
version_dct['files'] = [f.__dict__.copy() for f in version.files]
version_dct['images'] = [i.__dict__.copy() for i in version.images]
versions_dct.append(version_dct)
model_dct['versions'] = versions_dct
dct.append(model_dct)
# obj = json.dumps(dct, indent=2, ensure_ascii=False)
return dct
# ---------------------------------------------------------------------------
# Helpers
# ---------------------------------------------------------------------------
def model_to_dict(model) -> dict:
"""Convert a Pydantic CivitModel to a JSON-safe dict for v2 endpoints."""
return model.dict(by_alias=True)
def get_civitai(
model_id:int=None, # if model_id is provided assume fetch-from-cache
query:str = '', # search query or tag is required
tag:str = '', # search query or tag is required
types:str = '', # Checkpoint, TextualInversion, Hypernetwork, AestheticGradient, LORA, Controlnet, Poses
sort:str = '', # Highest Rated, Most Downloaded, Newest
period:str = '', # AllTime, Year, Month, Week, Day
nsfw:bool = None, # optional:bool
limit:int = 0,
base:str = '',
token:str = None,
exact:bool = True,
def version_to_dict(version) -> dict:
return version.dict(by_alias=True)
def model_to_legacy_dict(model) -> dict:
"""Convert a Pydantic CivitModel to the legacy v1 format expected by civitai.js."""
desc = model.description or ''
# Strip HTML tags for plain text desc
import re
text_desc = re.sub(r'<[^>]+>', '', desc).strip()
return {
'id': model.id,
'url': f'https://civitai.com/models/{model.id}',
'type': model.type,
'name': model.name,
'desc': text_desc[:200] if text_desc else '',
'html': desc,
'tags': model.tags,
'nsfw': model.nsfw,
'level': str(model.nsfw_level),
'availability': model.availability,
'downloads': model.stats.download_count,
'creator': model.creator.username if model.creator else 'Unknown',
'versions': [version_to_legacy_dict(v) for v in model.versions],
}
def version_to_legacy_dict(version) -> dict:
"""Convert a Pydantic CivitVersion to the legacy v1 format expected by civitai.js."""
desc = version.description or ''
import re
text_desc = re.sub(r'<[^>]+>', '', desc).strip()
return {
'id': version.id,
'name': version.name,
'base': version.base_model,
'mtime': version.published_at or '',
'downloads': version.stats.download_count,
'availability': version.availability,
'desc': text_desc[:200] if text_desc else '',
'html': desc,
'files': [file_to_legacy_dict(f) for f in version.files],
'images': [{'id': i.id, 'url': i.url, 'width': i.width, 'height': i.height, 'type': i.type} for i in version.images],
}
def file_to_legacy_dict(f) -> dict:
"""Convert a Pydantic CivitFile to the legacy v1 format expected by civitai.js."""
return {
'id': f.id,
'size': int(f.size_kb * 1024),
'name': f.name,
'type': f.type,
'hashes': [h for h in [f.hashes.sha256, f.hashes.autov1, f.hashes.autov2, f.hashes.autov3, f.hashes.crc32, f.hashes.blake3] if h],
'url': f.download_url,
}
# ---------------------------------------------------------------------------
# Search & Lookup
# ---------------------------------------------------------------------------
def get_search(
query: str = '',
tag: str = '',
types: str = '',
sort: str = '',
period: str = '',
base_models: str = '',
nsfw: bool = None,
limit: int = 20,
cursor: str = None,
username: str = '',
favorites: bool = False,
token: str = None,
):
"""Search CivitAI models with pagination."""
from modules.civitai.client_civitai import client
from modules.civitai.userdata_civitai import search_history
bm_list = [b.strip() for b in base_models.split(',') if b.strip()] if base_models else None
response = client.search_models(
query=query, tag=tag, types=types, sort=sort, period=period,
base_models=bm_list, nsfw=nsfw, limit=limit, cursor=cursor,
username=username, favorites=favorites, token=token,
)
if query:
search_history.add('query', query)
elif tag:
search_history.add('tag', tag)
return response.dict(by_alias=True)
def get_model(model_id: int, token: str = None):
"""Get a single model by ID (fresh fetch)."""
from modules.civitai.client_civitai import client
model = client.get_model(model_id, token=token)
if model is None:
return JSONResponse(content={"error": "model not found"}, status_code=404)
return model_to_dict(model)
def get_version(version_id: int, token: str = None):
"""Get a single version by ID."""
from modules.civitai.client_civitai import client
version = client.get_version(version_id, token=token)
if version is None:
return JSONResponse(content={"error": "version not found"}, status_code=404)
return version_to_dict(version)
def get_version_by_hash(hash_str: str, token: str = None):
"""Look up a version by file hash."""
from modules.civitai.client_civitai import client
version = client.get_version_by_hash(hash_str, token=token)
if version is None:
return JSONResponse(content={"error": "version not found"}, status_code=404)
return version_to_dict(version)
def get_options():
"""Get valid types, sort, period, base_models from CivitAI API discovery."""
from modules.civitai.client_civitai import client
return client.discover_options()
def get_tags(query: str = '', limit: int = 20, page: int = 1):
"""Search CivitAI tags."""
from modules.civitai.client_civitai import client
return client.get_tags(query=query, limit=limit, page=page).dict(by_alias=True)
def get_creators(query: str = '', limit: int = 20, page: int = 1):
"""Search CivitAI creators."""
from modules.civitai.client_civitai import client
return client.get_creators(query=query, limit=limit, page=page).dict(by_alias=True)
def get_images(model_id: int = None, model_version_id: int = None, limit: int = 20):
"""Get images with generation metadata from CivitAI."""
from modules.civitai.client_civitai import client
return {"items": client.get_images_raw(model_id=model_id, model_version_id=model_version_id, limit=limit)}
def get_me(token: str = None):
"""Get authenticated CivitAI user profile."""
from modules.civitai.client_civitai import client
profile = client.get_me(token=token)
if profile is None:
return JSONResponse(content={"error": "not authenticated"}, status_code=401)
return profile.dict(by_alias=True)
# ---------------------------------------------------------------------------
# Download Queue
# ---------------------------------------------------------------------------
def post_download(request: dict):
"""Queue a download. Returns the download item with its ID."""
from modules.civitai.download_civitai import download_manager
url = request.get('url', '')
filename = request.get('filename', '')
folder = request.get('folder', '')
model_type = request.get('model_type', 'Checkpoint')
expected_hash = request.get('expected_hash', '')
token = request.get('token', None)
model_name = request.get('model_name', '')
base_model = request.get('base_model', '')
creator = request.get('creator', '')
model_id = request.get('model_id', 0)
version_id = request.get('version_id', 0)
version_name = request.get('version_name', '')
nsfw = request.get('nsfw', False)
if not url:
return JSONResponse(content={"error": "url is required"}, status_code=400)
if not folder:
from modules.civitai.filemanage_civitai import resolve_save_path
folder = str(resolve_save_path(
model_type, model_name=model_name, base_model=base_model,
nsfw=nsfw, creator=creator, model_id=model_id,
version_id=version_id, version_name=version_name,
))
else:
if not os.path.isabs(folder):
from modules import paths
folder = os.path.join(paths.models_path, folder)
from modules import paths as _paths
if not os.path.realpath(folder).startswith(os.path.realpath(_paths.models_path)):
return JSONResponse(content={"error": "path outside models directory"}, status_code=400)
item = download_manager.enqueue(
url=url,
folder=folder,
filename=filename or "Unknown",
model_type=model_type,
expected_hash=expected_hash,
token=token,
model_id=int(model_id) if model_id else 0,
version_id=int(version_id) if version_id else 0,
)
return item.to_dict()
def post_download_cancel(download_id: str):
"""Cancel a queued or active download."""
from modules.civitai.download_civitai import download_manager
result = download_manager.cancel(download_id)
if not result:
return JSONResponse(content={"error": "download not found or already completed"}, status_code=404)
return {"success": True, "id": download_id}
def get_download_status():
"""Get the full download queue status."""
from modules.civitai.download_civitai import download_manager
return download_manager.status()
# ---------------------------------------------------------------------------
# Settings
# ---------------------------------------------------------------------------
def get_settings():
"""Get CivitAI-related settings."""
from modules import shared
return {
"token_configured": bool(getattr(shared.opts, 'civitai_token', '')),
"save_subfolder_enabled": getattr(shared.opts, 'civitai_save_subfolder_enabled', False),
"save_subfolder": getattr(shared.opts, 'civitai_save_subfolder', '{{BASEMODEL}}'),
"save_type_folders": getattr(shared.opts, 'civitai_save_type_folders', ''),
"discard_hash_mismatch": getattr(shared.opts, 'civitai_discard_hash_mismatch', True),
"download_workers": getattr(shared.opts, 'civitai_download_workers', 2),
}
def post_settings(request: dict):
"""Update CivitAI-related settings and persist to config."""
from modules import shared
token = request.get('token')
save_subfolder = request.get('save_subfolder')
discard_hash_mismatch = request.get('discard_hash_mismatch')
if token is not None:
if token.strip():
from modules.civitai.client_civitai import client
user = client.validate_token(token.strip())
if user is None:
return JSONResponse(content={"error": "Invalid API token"}, status_code=400)
log.info(f'CivitAI token validated: user={user.get("username", "?")}')
shared.opts.data['civitai_token'] = token.strip()
save_subfolder_enabled = request.get('save_subfolder_enabled')
if save_subfolder_enabled is not None:
shared.opts.data['civitai_save_subfolder_enabled'] = bool(save_subfolder_enabled)
if save_subfolder is not None:
shared.opts.data['civitai_save_subfolder'] = save_subfolder
if discard_hash_mismatch is not None:
shared.opts.data['civitai_discard_hash_mismatch'] = discard_hash_mismatch
shared.opts.save(shared.config_filename)
return get_settings()
def get_resolve_path(
model_type: str = 'Checkpoint',
model_name: str = '',
base_model: str = '',
creator: str = '',
model_id: int = 0,
version_id: int = 0,
version_name: str = '',
nsfw: bool = False,
):
"""Preview where a download would be saved."""
from modules.civitai.filemanage_civitai import resolve_save_path
path = resolve_save_path(
model_type, model_name=model_name, base_model=base_model,
nsfw=nsfw, creator=creator, model_id=model_id,
version_id=version_id, version_name=version_name,
)
return {"path": str(path)}
# ---------------------------------------------------------------------------
# Metadata
# ---------------------------------------------------------------------------
def post_metadata_scan(request: dict = None):
"""Scan local models for CivitAI metadata. Optional ``page`` filters by network type (e.g. 'lora', 'model')."""
from modules.civitai import metadata_civitai
page = (request or {}).get('page', None)
results = []
for batch in metadata_civitai.civit_search_metadata(title=page, raw=True):
if isinstance(batch, list):
results = batch
return {"results": results}
def post_metadata_update():
"""Update local metadata from CivitAI."""
from modules.civitai import metadata_civitai
items = []
for batch in metadata_civitai.civit_update_metadata(raw=True):
if isinstance(batch, list):
items = batch
results = []
for item in items:
results.append({
"file": getattr(item, "file", None),
"id": getattr(item, "id", None),
"name": getattr(item, "name", None),
"sha": getattr(item, "sha", None),
"versions": getattr(item, "versions", None),
"latest": getattr(item, "latest_name", None),
"status": getattr(item, "status", None),
})
return {"results": results}
# ---------------------------------------------------------------------------
# User Data — Bookmarks
# ---------------------------------------------------------------------------
def get_bookmarks():
from modules.civitai.userdata_civitai import bookmarks
return {"bookmarks": [{"name": n} for n in bookmarks.list()]}
def post_bookmark(request: dict):
from modules.civitai.userdata_civitai import bookmarks
name = request.get('name', '')
if not name:
return JSONResponse(content={"error": "name is required"}, status_code=400)
added = bookmarks.add(name)
return {"success": added, "name": name}
def delete_bookmark(name: str):
from modules.civitai.userdata_civitai import bookmarks
removed = bookmarks.remove(name)
if not removed:
return JSONResponse(content={"error": "not found"}, status_code=404)
return {"success": True, "name": name}
# ---------------------------------------------------------------------------
# User Data — Banned
# ---------------------------------------------------------------------------
def get_banned():
from modules.civitai.userdata_civitai import banned
return {"banned": [{"name": n} for n in banned.list()]}
def post_banned(request: dict):
from modules.civitai.userdata_civitai import banned
name = request.get('name', '')
if not name:
return JSONResponse(content={"error": "name is required"}, status_code=400)
added = banned.add(name)
return {"success": added, "name": name}
def delete_banned(name: str):
from modules.civitai.userdata_civitai import banned
removed = banned.remove(name)
if not removed:
return JSONResponse(content={"error": "not found"}, status_code=404)
return {"success": True, "name": name}
# ---------------------------------------------------------------------------
# User Data — Search History
# ---------------------------------------------------------------------------
def get_history(search_type: str = None):
from modules.civitai.userdata_civitai import search_history
return {"history": search_history.list(search_type)}
def delete_history():
from modules.civitai.userdata_civitai import search_history
search_history.clear()
return {"success": True}
# ---------------------------------------------------------------------------
# Sidecar Index (lazy-built from CivitAI .json metadata files)
# ---------------------------------------------------------------------------
sidecar_index = None
def buildsidecar_index():
"""Scan CivitAI .json sidecar files to build sha256 -> local file mapping. Lazy, one-time."""
global sidecar_index # pylint: disable=global-statement
if sidecar_index is not None:
return sidecar_index
sidecar_index = {}
import glob
from modules import shared, paths
model_exts = ('.safetensors', '.ckpt', '.pt', '.pth', '.bin')
dirs = [
(getattr(shared.cmd_opts, 'lora_dir', None) or os.path.join(paths.models_path, 'Lora'), 'lora'),
(getattr(shared.opts, 'ckpt_dir', None) or os.path.join(paths.models_path, 'Stable-diffusion'), 'checkpoint'),
]
for model_dir, model_type in dirs:
if not os.path.isdir(model_dir):
continue
for json_path in glob.glob(os.path.join(model_dir, '**', '*.json'), recursive=True):
try:
from modules.json_helpers import readfile
data = readfile(json_path, silent=True, as_type="dict")
if not isinstance(data, dict) or 'modelVersions' not in data:
continue
# Find the companion model file for this sidecar
base = os.path.splitext(json_path)[0]
companion = None
for ext in model_exts:
candidate = base + ext
if os.path.isfile(candidate):
companion = candidate
break
if not companion:
continue
# Match the companion file to a JSON entry by size (sizeKB)
companion_size_kb = os.path.getsize(companion) / 1024.0
companion_name = os.path.basename(base)
best_sha = None
best_diff = float('inf')
for v in data.get('modelVersions', []):
for f in v.get('files', []):
sha = (f.get('hashes') or {}).get('SHA256')
size_kb = f.get('sizeKB', 0)
if sha and size_kb and abs(size_kb - companion_size_kb) < 1.0:
diff = abs(size_kb - companion_size_kb)
if diff < best_diff:
best_sha = sha
best_diff = diff
if best_sha:
sidecar_index[best_sha.lower()] = {"filename": companion_name, "type": model_type}
except Exception:
continue
log.debug(f'CivitAI sidecar index: {len(sidecar_index)} hashes from sidecar files')
return sidecar_index
def invalidatesidecar_index():
global sidecar_index # pylint: disable=global-statement
sidecar_index = None
# ---------------------------------------------------------------------------
# Local Hash Check
# ---------------------------------------------------------------------------
def post_check_local(request: dict):
"""Check which SHA256 hashes correspond to locally downloaded files."""
from modules import hashes as hash_module
input_hashes = request.get('hashes', [])
if not input_hashes:
return {"found": {}}
# Build reverse lookup: lowercase sha256 -> {filename, type}
found = {}
hash_cache = hash_module.cache("hashes")
for title, entry in hash_cache.items():
sha = entry.get("sha256")
if not sha:
continue
parts = title.split("/", 1)
file_type = parts[0] if len(parts) > 1 else "unknown"
found[sha.lower()] = {"filename": title, "type": file_type}
# Supplement from in-memory checkpoint registry
try:
from modules.sd_checkpoint import checkpoints_list
for _title, cp in checkpoints_list.items():
if cp.sha256:
key = cp.sha256.lower()
if key not in found:
found[key] = {"filename": cp.filename, "type": "checkpoint"}
except Exception:
pass
# Supplement from in-memory LoRA registry
try:
from modules.lora.lora_load import available_networks
for _name, net in available_networks.items():
if net.hash:
key = net.hash.lower()
if key not in found:
found[key] = {"filename": net.filename, "type": "lora"}
except Exception:
pass
# Supplement from sidecar index (covers files never hashed locally)
sidecar = buildsidecar_index()
for h in input_hashes:
if not h:
continue
key = h.lower()
if key not in found and key in sidecar:
found[key] = sidecar[key]
# Match requested hashes
result = {}
for h in input_hashes:
if not h:
continue
match = found.get(h.lower())
if match:
result[h] = match
return {"found": result}
# ---------------------------------------------------------------------------
# Legacy Endpoints (backward compatibility)
# ---------------------------------------------------------------------------
def legacy_get_civitai(
model_id: int = None,
query: str = '',
tag: str = '',
types: str = '',
sort: str = '',
period: str = '',
nsfw: bool = None,
limit: int = 0,
base: str = '',
token: str = None,
exact: bool = True,
):
"""Legacy GET /sdapi/v1/civitai — delegates to search or model lookup."""
from modules.civitai import search_civitai
if model_id is not None:
dct = models_to_json(search_civitai.models, model_id=model_id)
return JSONResponse(content=dct, status_code=200)
if len(query) > 0 or len(tag) > 0:
# Legacy model_id lookup used the global cache; now we fetch fresh
from modules.civitai.client_civitai import client
model = client.get_model(model_id, token=token)
if model is None:
return JSONResponse(content=[], status_code=200)
return [model_to_legacy_dict(model)]
if query or tag:
models = search_civitai.search_civitai(
query=query,
tag=tag,
types=types,
sort=sort,
period=period,
nsfw=nsfw,
limit=limit,
base=base,
token=token,
exact=exact
query=query, tag=tag, types=types, sort=sort, period=period,
nsfw=nsfw, limit=limit, base=base, token=token, exact=exact,
)
dct = models_to_json(models)
return JSONResponse(content=dct, status_code=200)
return [model_to_legacy_dict(m) for m in models]
return JSONResponse(content=[], status_code=200)
def post_civitai(page:str=None):
def legacy_post_civitai(page: str = None):
"""Legacy POST /sdapi/v1/civitai — scan metadata."""
from modules.civitai import metadata_civitai
result = []
for r in metadata_civitai.civit_search_metadata(title=page, raw=True):
result = r # get the last yielded result
result = r
return result
# ---------------------------------------------------------------------------
# Registration
# ---------------------------------------------------------------------------
def register_api():
from modules.shared import api
api.add_api_route("/sdapi/v1/civitai", get_civitai, methods=["GET"], response_model=list)
api.add_api_route("/sdapi/v1/civitai", post_civitai, methods=["POST"], response_model=list)
# New REST API
api.add_api_route("/sdapi/v2/civitai/search", get_search, methods=["GET"], tags=["CivitAI"])
api.add_api_route("/sdapi/v2/civitai/model/{model_id}", get_model, methods=["GET"], tags=["CivitAI"])
api.add_api_route("/sdapi/v2/civitai/version/{version_id}", get_version, methods=["GET"], tags=["CivitAI"])
api.add_api_route("/sdapi/v2/civitai/version/by-hash/{hash_str}", get_version_by_hash, methods=["GET"], tags=["CivitAI"])
api.add_api_route("/sdapi/v2/civitai/options", get_options, methods=["GET"], tags=["CivitAI"])
api.add_api_route("/sdapi/v2/civitai/tags", get_tags, methods=["GET"], tags=["CivitAI"])
api.add_api_route("/sdapi/v2/civitai/creators", get_creators, methods=["GET"], tags=["CivitAI"])
api.add_api_route("/sdapi/v2/civitai/images", get_images, methods=["GET"], tags=["CivitAI"])
api.add_api_route("/sdapi/v2/civitai/me", get_me, methods=["GET"], tags=["CivitAI"])
api.add_api_route("/sdapi/v2/civitai/download", post_download, methods=["POST"], tags=["CivitAI"])
api.add_api_route("/sdapi/v2/civitai/download/{download_id}/cancel", post_download_cancel, methods=["POST"], tags=["CivitAI"])
api.add_api_route("/sdapi/v2/civitai/download/status", get_download_status, methods=["GET"], tags=["CivitAI"])
api.add_api_route("/sdapi/v2/civitai/settings", get_settings, methods=["GET"], tags=["CivitAI"])
api.add_api_route("/sdapi/v2/civitai/settings", post_settings, methods=["POST"], tags=["CivitAI"])
api.add_api_route("/sdapi/v2/civitai/resolve-path", get_resolve_path, methods=["GET"], tags=["CivitAI"])
api.add_api_route("/sdapi/v2/civitai/metadata/scan", post_metadata_scan, methods=["POST"], tags=["CivitAI"])
api.add_api_route("/sdapi/v2/civitai/metadata/update", post_metadata_update, methods=["POST"], tags=["CivitAI"])
api.add_api_route("/sdapi/v2/civitai/bookmarks", get_bookmarks, methods=["GET"], tags=["CivitAI"])
api.add_api_route("/sdapi/v2/civitai/bookmarks", post_bookmark, methods=["POST"], tags=["CivitAI"])
api.add_api_route("/sdapi/v2/civitai/bookmarks/{name}", delete_bookmark, methods=["DELETE"], tags=["CivitAI"])
api.add_api_route("/sdapi/v2/civitai/banned", get_banned, methods=["GET"], tags=["CivitAI"])
api.add_api_route("/sdapi/v2/civitai/banned", post_banned, methods=["POST"], tags=["CivitAI"])
api.add_api_route("/sdapi/v2/civitai/banned/{name}", delete_banned, methods=["DELETE"], tags=["CivitAI"])
api.add_api_route("/sdapi/v2/civitai/history", get_history, methods=["GET"], tags=["CivitAI"])
api.add_api_route("/sdapi/v2/civitai/history", delete_history, methods=["DELETE"], tags=["CivitAI"])
api.add_api_route("/sdapi/v2/civitai/check-local", post_check_local, methods=["POST"], tags=["CivitAI"])
# Legacy endpoints (backward compatibility)
api.add_api_route("/sdapi/v1/civitai", legacy_get_civitai, methods=["GET"], response_model=list, tags=["Models"])
api.add_api_route("/sdapi/v1/civitai", legacy_post_civitai, methods=["POST"], response_model=list, tags=["Models"])
log.debug('CivitAI API: registered endpoints')
+264
View File
@@ -0,0 +1,264 @@
import os
import time
from modules.logger import log
from modules.civitai.models_civitai import CivitModel, CivitVersion, CivitImage, CivitSearchResponse, CivitTagResponse, CivitCreatorResponse, CivitUserProfile
options_cache: dict = {}
options_cache_time: float = 0
OPTIONS_TTL = 3600 # 1 hour
class CivitaiClient:
BASE_URL = "https://civitai.com/api/v1"
def _get_token(self, token: str | None = None) -> str | None:
if token:
return token
from modules import shared
tok = getattr(shared.opts, 'civitai_token', '') or ''
if tok:
return tok
return os.environ.get('CIVITAI_TOKEN', None)
def _get(self, path: str, params: dict | None = None, token: str | None = None, stream: bool = False):
from modules import shared
url = f"{self.BASE_URL}{path}"
headers = {}
tok = self._get_token(token)
if tok:
headers['Authorization'] = f'Bearer {tok}'
if params:
from urllib.parse import urlencode
query = urlencode({k: v for k, v in params.items() if v is not None and v != ''}, doseq=True)
if query:
url = f"{url}?{query}"
return shared.req(url, headers=headers if headers else None, stream=stream)
def search_models(self, *, query: str = "", tag: str = "", types: str = "", sort: str = "", period: str = "",
base_models: list[str] | None = None, nsfw: bool | None = None, limit: int = 20,
cursor: str | None = None, username: str = "", favorites: bool = False,
token: str | None = None) -> CivitSearchResponse:
params: dict = {}
if query:
params['query'] = query
if tag:
params['tag'] = tag
if types:
params['types'] = types
if sort:
params['sort'] = sort
if period:
params['period'] = period
if base_models:
params['baseModels'] = base_models
if nsfw is not None:
params['nsfw'] = 'true' if nsfw else 'false'
if limit:
params['limit'] = limit
if cursor:
params['cursor'] = cursor
if username:
params['username'] = username
if favorites:
params['favorites'] = 'true'
r = self._get('/models', params=params, token=token)
if r.status_code != 200:
log.error(f'CivitAI search: code={r.status_code} reason={getattr(r, "reason", "")}')
return CivitSearchResponse()
data = r.json()
if 'items' not in data:
# single model by numeric query — wrap in search response
try:
model = CivitModel.parse_obj(data)
return CivitSearchResponse(items=[model])
except Exception:
return CivitSearchResponse()
try:
return CivitSearchResponse.parse_obj(data)
except Exception as e:
log.error(f'CivitAI search parse error: {e}')
return CivitSearchResponse()
def get_model(self, model_id: int, *, token: str | None = None) -> CivitModel | None:
r = self._get(f'/models/{model_id}', token=token)
if r.status_code != 200:
log.error(f'CivitAI get model: id={model_id} code={r.status_code}')
return None
try:
return CivitModel.parse_obj(r.json())
except Exception as e:
log.error(f'CivitAI get model parse error: id={model_id} {e}')
return None
def get_version(self, version_id: int, *, token: str | None = None) -> CivitVersion | None:
r = self._get(f'/model-versions/{version_id}', token=token)
if r.status_code != 200:
log.error(f'CivitAI get version: id={version_id} code={r.status_code}')
return None
try:
return CivitVersion.parse_obj(r.json())
except Exception as e:
log.error(f'CivitAI get version parse error: id={version_id} {e}')
return None
def get_version_by_hash(self, hash_str: str, *, token: str | None = None) -> CivitVersion | None:
r = self._get(f'/model-versions/by-hash/{hash_str}', token=token)
if r.status_code != 200:
return None
try:
return CivitVersion.parse_obj(r.json())
except Exception as e:
log.error(f'CivitAI get version by hash parse error: hash={hash_str} {e}')
return None
def get_images(self, *, model_version_id: int | None = None, limit: int | None = None, token: str | None = None) -> list[CivitImage]:
params: dict = {}
if model_version_id is not None:
params['modelVersionId'] = model_version_id
if limit is not None:
params['limit'] = limit
r = self._get('/images', params=params, token=token)
if r.status_code != 200:
return []
data = r.json()
items = data.get('items', [])
result = []
for item in items:
try:
result.append(CivitImage.parse_obj(item))
except Exception:
pass
return result
def get_images_raw(self, *, model_version_id: int | None = None, model_id: int | None = None, limit: int | None = None, token: str | None = None) -> list[dict]:
params: dict = {}
if model_version_id is not None:
params['modelVersionId'] = model_version_id
if model_id is not None:
params['modelId'] = model_id
if limit is not None:
params['limit'] = limit
r = self._get('/images', params=params, token=token)
if r.status_code != 200:
return []
data = r.json()
return data.get('items', [])
def get_tags(self, *, query: str = "", limit: int = 20, page: int = 1) -> CivitTagResponse:
params: dict = {}
if query:
params['query'] = query
if limit:
params['limit'] = limit
if page > 1:
params['page'] = page
r = self._get('/tags', params=params)
if r.status_code != 200:
return CivitTagResponse()
try:
return CivitTagResponse.parse_obj(r.json())
except Exception as e:
log.error(f'CivitAI get tags parse error: {e}')
return CivitTagResponse()
def get_creators(self, *, query: str = "", limit: int = 20, page: int = 1) -> CivitCreatorResponse:
params: dict = {}
if query:
params['query'] = query
if limit:
params['limit'] = limit
if page > 1:
params['page'] = page
r = self._get('/creators', params=params)
if r.status_code != 200:
return CivitCreatorResponse()
try:
return CivitCreatorResponse.parse_obj(r.json())
except Exception as e:
log.error(f'CivitAI get creators parse error: {e}')
return CivitCreatorResponse()
def get_me(self, token: str | None = None) -> CivitUserProfile | None:
r = self._get('/me', token=token)
if r.status_code != 200:
return None
try:
return CivitUserProfile.parse_obj(r.json())
except Exception as e:
log.error(f'CivitAI get me parse error: {e}')
return None
def validate_token(self, token: str) -> dict | None:
"""Validate a token by calling /me. Returns user info dict or None if invalid."""
profile = self.get_me(token=token)
if profile is None:
return None
return {"username": profile.username, "id": profile.id}
def discover_options(self) -> dict:
global options_cache, options_cache_time # pylint: disable=global-statement
now = time.time()
if options_cache and (now - options_cache_time) < OPTIONS_TTL:
return options_cache
from modules import shared
result: dict = {'types': [], 'sort': [], 'period': [], 'base_models': []}
# Send invalid params to trigger 400 with valid enum values in error response
probes = [
('types', '/models', {'types': '__invalid__'}),
('sort', '/models', {'sort': '__invalid__'}),
('period', '/models', {'period': '__invalid__'}),
('base_models', '/models', {'baseModels': '__invalid__'}),
]
for key, path, params in probes:
try:
url = f"{self.BASE_URL}{path}"
from urllib.parse import urlencode
query = urlencode(params)
full_url = f"{url}?{query}"
r = shared.req(full_url)
if r.status_code == 400:
data = r.json()
error = data.get('error', {})
if not isinstance(error, dict):
continue
# Parse ZodError: error.message is a JSON-encoded array of issues
import json as _json
issues = error.get('issues', [])
if not issues:
try:
issues = _json.loads(error.get('message', '[]'))
except Exception:
issues = []
for issue in issues:
# Flat format: options directly on issue
options = issue.get('options', [])
if options:
result[key] = options
break
# Flat format: values directly on issue (sort/period use this)
values = issue.get('values', [])
if values:
result[key] = values
break
# Nested union format: errors[][].values
for err_group in issue.get('errors', []):
if isinstance(err_group, list):
for err in err_group:
vals = err.get('values', [])
if vals:
result[key] = vals
break
if result[key]:
break
if result[key]:
break
except Exception as e:
log.debug(f'CivitAI discover options: key={key} {e}')
options_cache = result
options_cache_time = now
log.debug(f'CivitAI options: types={len(result["types"])} sort={len(result["sort"])} period={len(result["period"])} base_models={len(result["base_models"])}')
return result
client = CivitaiClient()
+380 -139
View File
@@ -1,31 +1,350 @@
import os
import json
import uuid
import hashlib
import threading
import time
from collections import deque
from dataclasses import dataclass, field
from datetime import datetime
import rich.progress as p
from PIL import Image
from modules import shared, errors, paths
from modules.logger import log, console
from modules.json_helpers import writefile
from installer import log
from modules import shared, paths
from modules.logger import console
pbar = None
@dataclass
class DownloadItem:
id: str = field(default_factory=lambda: uuid.uuid4().hex[:12])
url: str = ""
filename: str = ""
folder: str = ""
model_type: str = ""
expected_hash: str = ""
token: str | None = None
model_id: int = 0
version_id: int = 0
status: str = "queued" # queued | downloading | verifying | completed | failed | cancelled
progress: float = 0.0
bytes_downloaded: int = 0
bytes_total: int = 0
error: str | None = None
created_at: datetime = field(default_factory=datetime.now)
completed_at: datetime | None = None
def to_dict(self) -> dict:
return {
"id": self.id,
"url": self.url,
"filename": self.filename,
"folder": self.folder,
"model_type": self.model_type,
"status": self.status,
"progress": round(self.progress, 4),
"bytes_downloaded": self.bytes_downloaded,
"bytes_total": self.bytes_total,
"error": self.error,
"created_at": self.created_at.isoformat(),
"completed_at": self.completed_at.isoformat() if self.completed_at else None,
}
def save_video_frame(filepath: str):
from modules import video
try:
frames, fps, duration, w, h, codec, frame = video.get_video_params(filepath, capture=True)
except Exception as e:
log.error(f'Video: file={filepath} {e}')
return None
if frame is not None:
basename = os.path.splitext(filepath)
thumb = f'{basename[0]}.thumb.jpg'
log.debug(f'Video: file={filepath} frames={frames} fps={fps} size={w}x{h} codec={codec} duration={duration} thumb={thumb}')
frame.save(thumb)
else:
log.error(f'Video: file={filepath} no frames found')
return frame
class DownloadManager:
def __init__(self, max_workers: int = 2):
self._queue: deque[DownloadItem] = deque()
self._active: dict[str, DownloadItem] = {}
self._completed: deque[DownloadItem] = deque(maxlen=50)
self._cancel_ids: set[str] = set()
self._lock = threading.Lock()
self._max_workers = max_workers
self._worker_count = 0
def enqueue(self, url: str, folder: str, filename: str, model_type: str = "",
expected_hash: str = "", token: str | None = None,
model_id: int = 0, version_id: int = 0) -> DownloadItem:
item = DownloadItem(
url=url,
filename=filename,
folder=folder,
model_type=model_type,
expected_hash=expected_hash,
token=token,
model_id=model_id,
version_id=version_id,
)
with self._lock:
self._queue.append(item)
log.info(f'CivitAI download queued: id={item.id} file="{filename}" url="{url}"')
self._try_start_worker()
return item
def cancel(self, download_id: str) -> bool:
with self._lock:
# Check if in queue (not yet started)
for item in self._queue:
if item.id == download_id:
item.status = "cancelled"
item.completed_at = datetime.now()
self._queue.remove(item)
self._completed.append(item)
log.info(f'CivitAI download cancelled (queued): id={download_id}')
return True
# Check if active
if download_id in self._active:
self._cancel_ids.add(download_id)
log.info(f'CivitAI download cancel requested: id={download_id}')
return True
return False
def status(self) -> dict:
with self._lock:
return {
"active": [item.to_dict() for item in self._active.values()],
"queued": [item.to_dict() for item in self._queue],
"completed": [item.to_dict() for item in self._completed],
}
def get_active_items(self) -> list[dict]:
with self._lock:
return [item.to_dict() for item in self._active.values()]
def _try_start_worker(self):
with self._lock:
if self._worker_count >= self._max_workers:
return
if not self._queue:
return
self._worker_count += 1
thread = threading.Thread(target=self._worker, daemon=True)
thread.start()
def _worker(self):
try:
while True:
item = None
with self._lock:
if not self._queue:
break
item = self._queue.popleft()
self._active[item.id] = item
if item:
self._download(item)
with self._lock:
self._active.pop(item.id, None)
self._completed.append(item)
finally:
with self._lock:
self._worker_count -= 1
# Start more workers if items still queued
self._try_start_worker()
def _download(self, item: DownloadItem):
# Create temp file name from URL hash
url_hash = hashlib.sha256(item.url.encode('utf-8')).hexdigest()[:8]
temp_file = os.path.join(item.folder, f'{url_hash}.tmp')
final_file = os.path.join(item.folder, item.filename)
# Check if already exists
if os.path.isfile(final_file):
item.status = "completed"
item.progress = 1.0
item.completed_at = datetime.now()
item.error = "already exists"
log.info(f'CivitAI download: id={item.id} file="{final_file}" already exists')
return
# Ensure folder exists
os.makedirs(item.folder, exist_ok=True)
# Resume support
headers = {}
starting_pos = 0
if os.path.isfile(temp_file):
starting_pos = os.path.getsize(temp_file)
headers['Range'] = f'bytes={starting_pos}-'
# Auth
token = item.token or self._get_token()
if token and 'civit' in item.url.lower():
headers['Authorization'] = f'Bearer {token}'
item.status = "downloading"
item.bytes_downloaded = starting_pos
try:
r = shared.req(item.url, headers=headers if headers else None, stream=True)
if r.status_code not in (200, 206):
item.status = "failed"
item.error = f'HTTP {r.status_code}'
item.completed_at = datetime.now()
return
total_size = int(r.headers.get('content-length', 0))
item.bytes_total = starting_pos + total_size
# Resolve filename from Content-Disposition if not set
if not item.filename or item.filename == "Unknown":
cn = r.headers.get('content-disposition', '')
if 'filename=' in cn:
item.filename = cn.split('filename=')[-1].strip('"')
final_file = os.path.join(item.folder, item.filename)
block_size = 65536 # 64KB blocks
written = starting_pos
log.info(f'CivitAI download: id={item.id} file="{item.filename}" size={round((starting_pos + total_size) / 1024 / 1024, 1)}MB')
pbar = p.Progress(p.TextColumn('[cyan]{task.description}'), p.DownloadColumn(), p.BarColumn(), p.TaskProgressColumn(), p.TimeRemainingColumn(), p.TimeElapsedColumn(), p.TransferSpeedColumn(), p.TextColumn('[cyan]{task.fields[name]}'), console=console)
with pbar:
task = pbar.add_task(description="Download", total=starting_pos + total_size, name=item.filename)
with open(temp_file, 'ab') as f:
for chunk in r.iter_content(block_size):
# Check cancellation
if item.id in self._cancel_ids:
with self._lock:
self._cancel_ids.discard(item.id)
item.status = "cancelled"
item.completed_at = datetime.now()
log.info(f'CivitAI download cancelled: id={item.id}')
try:
os.remove(temp_file)
except OSError:
pass
return
f.write(chunk)
written += len(chunk)
item.bytes_downloaded = written
if item.bytes_total > 0:
item.progress = written / item.bytes_total
pbar.update(task, completed=written)
# Validate minimum size
if written < 1024:
try:
os.remove(temp_file)
except OSError:
pass
item.status = "failed"
item.error = f'download too small: {written} bytes'
item.completed_at = datetime.now()
return
# Check for incomplete download
if starting_pos + total_size != written:
item.status = "failed"
item.error = f'incomplete: expected={starting_pos + total_size} got={written}'
item.completed_at = datetime.now()
return
except Exception as e:
item.status = "failed"
item.error = str(e)
item.completed_at = datetime.now()
log.error(f'CivitAI download error: id={item.id} {e}')
return
# Hash verification
if item.expected_hash:
item.status = "verifying"
try:
from modules import hashes
computed = hashes.calculate_sha256(temp_file, quiet=True)
if computed.upper() != item.expected_hash.upper():
discard = getattr(shared.opts, 'civitai_discard_hash_mismatch', True)
if discard:
try:
os.remove(temp_file)
except OSError:
pass
item.status = "failed"
item.error = f'hash mismatch: expected={item.expected_hash[:16]}... got={computed[:16]}...'
item.completed_at = datetime.now()
log.error(f'CivitAI download hash mismatch: id={item.id} expected={item.expected_hash[:16]} got={computed[:16]}')
return
log.warning(f'CivitAI download hash mismatch (kept): id={item.id} expected={item.expected_hash[:16]} got={computed[:16]}')
except Exception as e:
log.warning(f'CivitAI download hash check failed: id={item.id} {e}')
# Move temp to final
try:
os.rename(temp_file, final_file)
except OSError as e:
item.status = "failed"
item.error = f'rename failed: {e}'
item.completed_at = datetime.now()
return
item.status = "completed"
item.progress = 1.0
item.completed_at = datetime.now()
log.info(f'CivitAI download complete: id={item.id} file="{final_file}" size={item.bytes_downloaded}')
# Write verified hash to cache so check-local finds it immediately
if item.expected_hash:
try:
from modules import hashes
model_type_map = {'Checkpoint': 'checkpoint', 'LORA': 'lora', 'TextualInversion': 'embedding', 'VAE': 'vae'}
prefix = model_type_map.get(item.model_type, item.model_type.lower())
name = os.path.splitext(item.filename)[0]
title = f"{prefix}/{name}"
hash_cache = hashes.cache("hashes")
hash_cache[title] = {
"mtime": os.path.getmtime(final_file),
"sha256": item.expected_hash.lower(),
}
hashes.dump_cache()
except Exception:
pass
# Download metadata and preview
self._fetch_sidecar(item, final_file)
# Refresh model list and extra-networks cache
try:
from modules.sd_models import list_models
list_models()
except Exception:
pass
try:
from modules.api.loras import _invalidate_extra_networks
_invalidate_extra_networks()
except Exception:
pass
def _fetch_sidecar(self, item: DownloadItem, final_file: str):
"""Download metadata JSON and preview image for a completed download."""
if not item.model_id:
return
try:
code, _size, _note = download_civit_meta(final_file, item.model_id)
if code == 200:
log.info(f'CivitAI metadata saved: id={item.id} model_id={item.model_id}')
except Exception as e:
log.warning(f'CivitAI metadata fetch failed: id={item.id} {e}')
if not item.version_id:
return
try:
from modules.civitai.client_civitai import client
version = client.get_version(item.version_id)
if version and version.images:
for img in version.images:
if img.url:
code, _size, _note = download_civit_preview(final_file, img.url)
if code == 200:
log.info(f'CivitAI preview saved: id={item.id}')
break
except Exception as e:
log.warning(f'CivitAI preview fetch failed: id={item.id} {e}')
def _get_token(self) -> str | None:
tok = getattr(shared.opts, 'civitai_token', '') or ''
if tok:
return tok
return os.environ.get('CIVITAI_TOKEN', None)
download_manager = DownloadManager()
# ---- Legacy compatibility functions ----
def download_civit_meta(model_path: str, model_id):
fn = os.path.splitext(model_path)[0] + '.json'
@@ -34,10 +353,12 @@ def download_civit_meta(model_path: str, model_id):
if r.status_code == 200:
try:
data = r.json()
from modules.json_helpers import writefile
writefile(data, filename=fn, mode='w', silent=True)
log.info(f'CivitAI download: id={model_id} url={url} file="{fn}"')
return r.status_code, len(data), '' # code/size/note
return r.status_code, len(data), ''
except Exception as e:
from modules import errors
errors.display(e, 'civitai meta')
log.error(f'CivitAI meta: id={model_id} url={url} file="{fn}" {e}')
return r.status_code, '', str(e)
@@ -45,9 +366,7 @@ def download_civit_meta(model_path: str, model_id):
def download_civit_preview(model_path: str, preview_url: str):
global pbar # pylint: disable=global-statement
if model_path is None:
pbar = None
return 500, '', ''
ext = os.path.splitext(preview_url)[1]
preview_file = os.path.splitext(model_path)[0] + ext
@@ -55,140 +374,62 @@ def download_civit_preview(model_path: str, preview_url: str):
is_json = preview_file.lower().endswith('.json')
if is_json:
log.warning(f'CivitAI download: url="{preview_url}" skip json')
return 500, '', 'exepected preview image got json'
return 500, '', 'expected preview image got json'
if os.path.exists(preview_file):
return 304, '', 'already exists'
# res = f'CivitAI download: url={preview_url} file="{preview_file}"'
r = shared.req(preview_url, stream=True)
total_size = int(r.headers.get('content-length', 0))
block_size = 16384 # 16KB blocks
block_size = 16384
written = 0
img = None
jobid = shared.state.begin('Download CivitAI')
if pbar is None:
pbar = p.Progress(p.TextColumn('[cyan]Download'), p.DownloadColumn(), p.BarColumn(), p.TaskProgressColumn(), p.TimeRemainingColumn(), p.TimeElapsedColumn(), p.TransferSpeedColumn(), p.TextColumn('[yellow]{task.description}'), console=console)
try:
with open(preview_file, 'wb') as f:
with pbar:
task = pbar.add_task(description=preview_file, total=total_size)
for data in r.iter_content(block_size):
written = written + len(data)
f.write(data)
pbar.update(task, advance=block_size)
if written < 1024: # min threshold
for data in r.iter_content(block_size):
written += len(data)
f.write(data)
if written < 1024:
os.remove(preview_file)
return 400, '', 'removed invalid download'
if is_video:
img = save_video_frame(preview_file)
from modules.civitai.video_helper import save_video_frame
save_video_frame(preview_file)
else:
from PIL import Image
img = Image.open(preview_file)
log.info(f'CivitAI download: url={preview_url} file="{preview_file}" size={total_size} image={img.size}')
img.close()
except Exception as e:
log.error(f'CivitAI download error: url={preview_url} file="{preview_file}" written={written} {e}')
shared.state.end(jobid)
return 500, '', str(e)
shared.state.end(jobid)
if img is None:
return 500, '', 'image is none'
log.info(f'CivitAI download: url={preview_url} file="{preview_file}" size={total_size} image={img.size}')
img.close()
return 200, str(total_size), '' # code/size/note
def download_civit_model_thread(model_name: str, model_url: str, model_path: str = "", model_type: str = "Model", token: str = None):
import hashlib
sha256 = hashlib.sha256()
sha256.update(model_url.encode('utf-8'))
temp_file = sha256.hexdigest()[:8] + '.tmp'
headers = {}
starting_pos = 0
if os.path.isfile(temp_file):
starting_pos = os.path.getsize(temp_file)
headers['Range'] = f'bytes={starting_pos}-'
if 'civit' in model_url.lower(): # downloader can be used for other urls too
if token is None or len(token) == 0:
token = shared.opts.civitai_token
if (token is not None) and (len(token) > 0):
headers['Authorization'] = f'Bearer {token}'
r = shared.req(model_url, headers=headers, stream=True)
total_size = int(r.headers.get('content-length', 0))
if model_name is None or len(model_name) == 0:
cn = r.headers.get('content-disposition', '')
model_name = cn.split('filename=')[-1].strip('"')
model_path = model_path.strip()
if len(model_path) > 0:
if os.path.isabs(model_path):
pass
else:
model_path = os.path.join(paths.models_path, model_path)
elif model_type.lower() == 'lora':
model_path = shared.opts.lora_dir
elif model_type.lower() == 'embedding':
model_path = shared.opts.embeddings_dir
elif model_type.lower() == 'vae':
model_path = shared.opts.vae_dir
else:
model_path = shared.opts.ckpt_dir
model_file = os.path.join(model_path, model_name)
temp_file = os.path.join(model_path, temp_file)
res = f'Model download: name="{model_name}" url="{model_url}" path="{model_path}" temp="{temp_file}"'
if os.path.isfile(model_file):
res += ' already exists'
log.warning(res)
return res
res += f' size={round((starting_pos + total_size)/1024/1024, 2)}Mb'
log.info(res)
jobid = shared.state.begin('Download CivitAI')
block_size = 16384 # 16KB blocks
written = starting_pos
global pbar # pylint: disable=global-statement
if pbar is None:
pbar = p.Progress(p.TextColumn('[cyan]{task.description}'), p.DownloadColumn(), p.BarColumn(), p.TaskProgressColumn(), p.TimeRemainingColumn(), p.TimeElapsedColumn(), p.TransferSpeedColumn(), p.TextColumn('[cyan]{task.fields[name]}'), console=console)
with pbar:
task = pbar.add_task(description="Download starting", total=starting_pos+total_size, name=model_name)
try:
with open(temp_file, 'ab') as f:
for data in r.iter_content(block_size):
if written == 0:
try: # check if response is JSON message instead of bytes
log.error(f'Model download: response={json.loads(data.decode("utf-8"))}')
raise ValueError('response: type=json expected=bytes')
except Exception: # this is good
pass
written = written + len(data)
f.write(data)
pbar.update(task, description="Download", completed=written)
if written < 1024: # min threshold
os.remove(temp_file)
raise ValueError(f'removed invalid download: bytes={written}')
except Exception as e:
log.error(f'{res} {e}')
finally:
pbar.stop_task(task)
pbar.remove_task(task)
if starting_pos+total_size != written:
log.warning(f'{res} written={round(written/1024/1024)}Mb incomplete download')
elif os.path.exists(temp_file):
log.debug(f'Model download complete: temp="{temp_file}" path="{model_file}"')
os.rename(temp_file, model_file)
shared.state.end(jobid)
if os.path.exists(model_file):
return model_file
else:
return None
return 200, str(total_size), ''
def download_civit_model(model_url: str, model_name: str = '', model_path: str = '', model_type: str = '', token: str = None):
import threading
if model_url is None or len(model_url) == 0:
"""Legacy function — delegates to DownloadManager for non-blocking downloads."""
if not model_url:
log.error('Model download: no url provided')
return
thread = threading.Thread(target=download_civit_model_thread, args=(model_name, model_url, model_path, model_type, token))
thread.start()
thread.join()
from modules.sd_models import list_models # pylint: disable=W0621
list_models()
return None
from modules.civitai.filemanage_civitai import get_type_folder
if not model_path:
folder = str(get_type_folder(model_type or 'Checkpoint'))
elif os.path.isabs(model_path):
folder = model_path
else:
folder = os.path.join(paths.models_path, model_path)
item = download_manager.enqueue(
url=model_url,
folder=folder,
filename=model_name or "Unknown",
model_type=model_type,
token=token,
)
# Wait for completion (legacy blocking behavior)
while item.status in ("queued", "downloading", "verifying"):
time.sleep(0.5)
if item.status == "completed" and not item.error:
from modules.sd_models import list_models
list_models()
return os.path.join(item.folder, item.filename)
return None
+97
View File
@@ -0,0 +1,97 @@
import os
import re
from pathlib import Path
from modules.logger import log
# Map CivitAI model types to shared.opts directory settings and fallback subfolder names
TYPE_MAP = {
'Checkpoint': ('ckpt_dir', 'Stable-diffusion'),
'TextualInversion': ('embeddings_dir', 'embeddings'),
'Hypernetwork': ('hypernetwork_dir', 'hypernetworks'),
'AestheticGradient': ('ckpt_dir', 'Stable-diffusion'),
'LORA': ('lora_dir', 'Lora'),
'LoCon': ('lora_dir', 'Lora'),
'DoRA': ('lora_dir', 'Lora'),
'Controlnet': ('control_dir', 'control'),
'Poses': ('ckpt_dir', 'Stable-diffusion'),
'Wildcards': ('wildcards_dir', 'wildcards'),
'Workflows': (None, 'workflows'),
'VAE': ('vae_dir', 'VAE'),
'MotionModule': (None, 'motion'),
'Upscaler': ('esrgan_models_path', 'ESRGAN'),
'Other': ('ckpt_dir', 'Stable-diffusion'),
}
def get_type_folder(model_type: str) -> Path:
from modules import shared, paths
# Check for user-configured type folder overrides
custom_json = getattr(shared.opts, 'civitai_save_type_folders', '') or ''
if custom_json.strip():
try:
import json
custom = json.loads(custom_json)
if model_type in custom:
p = Path(custom[model_type])
if p.is_absolute():
return p
return Path(paths.models_path) / custom[model_type]
except Exception as e:
log.warning(f'CivitAI type folder override parse error: {e}')
opt_attr, fallback_dir = TYPE_MAP.get(model_type, ('ckpt_dir', 'Stable-diffusion'))
if opt_attr:
configured = getattr(shared.opts, opt_attr, '') or ''
if configured:
return Path(configured)
return Path(paths.models_path) / fallback_dir
def resolve_save_path(model_type: str, model_name: str = "", base_model: str = "",
nsfw: bool = False, creator: str = "", model_id: int = 0,
version_id: int = 0, version_name: str = "") -> Path:
from modules import shared
base_folder = get_type_folder(model_type)
if not getattr(shared.opts, 'civitai_save_subfolder_enabled', False):
return base_folder
template = getattr(shared.opts, 'civitai_save_subfolder', '{{BASEMODEL}}') or ''
if not template:
return base_folder
# Template variable substitution
replacements = {
'{{BASEMODEL}}': sanitize_filename(base_model) if base_model else '_unknown',
'{{MODELNAME}}': sanitize_filename(model_name) if model_name else '',
'{{CREATOR}}': sanitize_filename(creator) if creator else '_unknown',
'{{MODELID}}': str(model_id) if model_id else '0',
'{{VERSIONID}}': str(version_id) if version_id else '0',
'{{VERSIONNAME}}': sanitize_filename(version_name) if version_name else '',
'{{NSFW}}': 'nsfw' if nsfw else 'sfw',
'{{TYPE}}': sanitize_filename(model_type) if model_type else 'other',
}
subfolder = template
for key, value in replacements.items():
subfolder = subfolder.replace(key, value)
# Clean up empty path segments
subfolder = re.sub(r'[/\\]+', os.sep, subfolder)
subfolder = subfolder.strip(os.sep)
return base_folder / subfolder
def check_exists(folder: Path, filename: str) -> bool:
return (folder / filename).exists()
def sanitize_filename(name: str) -> str:
if not name:
return ''
# Replace unsafe characters
name = re.sub(r'[<>:"/\\|?*\x00-\x1f]', '_', name)
# Collapse multiple underscores/spaces
name = re.sub(r'[_ ]{2,}', '_', name)
name = name.strip(' _.')
# Truncate to 200 chars (leaving room for extension and path)
if len(name.encode('utf-8')) > 200:
while len(name.encode('utf-8')) > 200:
name = name[:-1]
name = name.rstrip(' _.')
return name
+42 -99
View File
@@ -1,16 +1,12 @@
import os
import re
import time
import gradio as gr
from modules.shared import log, opts, req, readfile, max_workers
data = []
selected_model = None
from modules.shared import log, opts, readfile, max_workers
from modules.civitai.client_civitai import client
class CivitModel:
def __init__(self, name, fn, sha = None, meta = None):
def __init__(self, name, fn, sha=None, meta=None):
if meta is None:
meta = {}
self.name = name
@@ -28,7 +24,7 @@ class CivitModel:
self.status = 'Not found'
def civit_update_metadata(raw:bool=False):
def civit_update_metadata(raw: bool = False):
def create_update_metadata_table(rows: list[CivitModel]):
html = """
<table class="simple-table">
@@ -72,16 +68,15 @@ def civit_update_metadata(raw:bool=False):
if model.sha is None or len(model.sha) == 0:
log.debug(f'CivitAI skip search: name="{model.name}" hash=None')
else:
r = req(f'https://civitai.com/api/v1/model-versions/by-hash/{model.sha}')
log.debug(f'CivitAI search: name="{model.name}" hash={model.sha} status={r.status_code}')
if r.status_code == 200:
d = r.json()
model.id = d['modelId']
version = client.get_version_by_hash(model.sha)
if version is not None:
model.id = version.model_id
download_civit_meta(model.fn, model.id)
fn = os.path.splitext(item['filename'])[0] + '.json'
model.meta = readfile(fn, silent=True, as_type="dict")
model.name = model.meta.get('name', model.name)
model.versions = len(model.meta.get('modelVersions', []))
time.sleep(0.25) # rate limiting
versions = model.meta.get('modelVersions', [])
if len(versions) > 0:
model.latest = versions[0].get('name', '')
@@ -108,58 +103,6 @@ def civit_update_metadata(raw:bool=False):
yield results if raw else create_update_metadata_table(results)
def civit_search_model(name, tag, model_type):
# types = 'LORA' if model_type == 'LoRA' else 'Checkpoint'
url = 'https://civitai.com/api/v1/models?limit=25&Sort=Newest'
if model_type == 'Model':
url += '&types=Checkpoint'
elif model_type == 'LoRA':
url += '&types=LORA&types=DoRA&types=LoCon'
elif model_type == 'Embedding':
url += '&types=TextualInversion'
elif model_type == 'VAE':
url += '&types=VAE'
if name is not None and len(name) > 0:
url += f'&query={name}'
if tag is not None and len(tag) > 0:
url += f'&tag={tag}'
r = req(url)
log.debug(f'CivitAI search: type={model_type} name="{name}" tag={tag or "none"} url="{url}" status={r.status_code}')
if r.status_code != 200:
log.warning(f'CivitAI search: name="{name}" tag={tag} status={r.status_code}')
return [], gr.update(visible=False, value=[]), gr.update(visible=False, value=None), gr.update(visible=False, value=None)
try:
body = r.json()
except Exception as e:
log.error(f'CivitAI search: name="{name}" tag={tag} {e}')
return [], gr.update(visible=False, value=[]), gr.update(visible=False, value=None), gr.update(visible=False, value=None)
global data # pylint: disable=global-statement
data = body.get('items', [])
data1 = []
for model in data:
found = 0
if model_type == 'LoRA' and model['type'].lower() in ['lora', 'locon', 'dora', 'lycoris']:
found += 1
elif model_type == 'Embedding' and model['type'].lower() in ['textualinversion', 'embedding']:
found += 1
elif model_type == 'Model' and model['type'].lower() in ['checkpoint']:
found += 1
elif model_type == 'VAE' and model['type'].lower() in ['vae']:
found += 1
elif model_type == 'Other':
found += 1
if found > 0:
data1.append([
model['id'],
model['name'],
', '.join(model['tags']),
model['stats']['downloadCount'],
model['stats']['rating']
])
res = f'Search result: name={name} tag={tag or "none"} type={model_type} models={len(data1)}'
return res, gr.update(visible=len(data1) > 0, value=data1 if len(data1) > 0 else []), gr.update(visible=False, value=None), gr.update(visible=False, value=None)
def atomic_civit_search_metadata(item, results):
from modules.civitai.download_civitai import download_civit_preview, download_civit_meta
if item is None:
@@ -167,7 +110,6 @@ def atomic_civit_search_metadata(item, results):
try:
meta = os.path.splitext(item['filename'])[0] + '.json'
except Exception:
# log.error(f'CivitAI search metadata: item={item} {e}')
return
has_meta = os.path.isfile(meta) and os.stat(meta).st_size > 0
if ('missing.png' in item['preview'] or not has_meta) and os.path.isfile(item['filename']):
@@ -183,47 +125,48 @@ def atomic_civit_search_metadata(item, results):
'note': '',
}
if sha is not None and len(sha) > 0:
r = req(f'https://civitai.com/api/v1/model-versions/by-hash/{sha}')
log.debug(f'CivitAI search: name="{item["name"]}" hash={sha} status={r.status_code}')
version = client.get_version_by_hash(sha)
result['hash'] = sha
result['code'] = r.status_code
if r.status_code == 200:
d = r.json()
result['code'], result['size'], result['note'] = download_civit_meta(item['filename'], d['modelId'])
result['id'] = d['modelId']
if version is not None:
result['code'] = 200
result['code'], result['size'], result['note'] = download_civit_meta(item['filename'], version.model_id)
result['id'] = version.model_id
result['type'] = 'metadata'
results.append(result)
if d.get('images') is not None:
for i in d['images']:
result['code'], result['size'], result['note'] = download_civit_preview(item['filename'], i['url'])
if result['code'] == 200:
result['type'] = 'preview'
results.append(result)
# Create a new dict for each append to avoid mutation bugs
results.append(dict(result))
for img in version.images:
if img.url:
code, size, note = download_civit_preview(item['filename'], img.url)
if code == 200:
results.append({**result, 'code': code, 'size': size, 'note': note, 'type': 'preview'})
found = True
break
else:
result['code'] = 404
time.sleep(0.25) # rate limiting
if not found and os.stat(item['filename']).st_size < (1024 * 1024 * 1024):
from modules import hashes
sha = hashes.calculate_sha256(item['filename'], quiet=True)[:10]
r = req(f'https://civitai.com/api/v1/model-versions/by-hash/{sha}')
log.debug(f'CivitAI search: name="{item["name"]}" hash={sha} status={r.status_code}')
version = client.get_version_by_hash(sha)
result['hash'] = sha
result['code'] = r.status_code
if r.status_code == 200:
d = r.json()
result['code'], result['size'], result['note'] = download_civit_meta(item['filename'], d['modelId'])
result['id'] = d['modelId']
if version is not None:
result['code'] = 200
result['code'], result['size'], result['note'] = download_civit_meta(item['filename'], version.model_id)
result['id'] = version.model_id
result['type'] = 'metadata'
results.append(result)
if d.get('images') is not None:
for i in d['images']:
result['code'], result['size'], result['note'] = download_civit_preview(item['filename'], i['url'])
if result['code'] == 200:
result['type'] = 'preview'
results.append(result)
results.append(dict(result))
for img in version.images:
if img.url:
code, size, note = download_civit_preview(item['filename'], img.url)
if code == 200:
results.append({**result, 'code': code, 'size': size, 'note': note, 'type': 'preview'})
found = True
break
else:
result['code'] = 404
time.sleep(0.25) # rate limiting
if not found:
results.append(result)
results.append(dict(result))
def civit_search_metadata(title: str = None, raw: bool = False):
@@ -259,10 +202,10 @@ def civit_search_metadata(title: str = None, raw: bool = False):
candidates = []
re_skip = [r.strip() for r in opts.extra_networks_scan_skip.split(',') if len(r.strip()) > 0]
for page in get_pages():
if type(title) == str:
if isinstance(title, str):
if page.title.lower() != title.lower():
continue
if page.name == 'style' or page.name == 'wildcards':
if page.name in ('style', 'wildcards'):
continue
for item in page.list_items():
if item is None:
@@ -272,7 +215,7 @@ def civit_search_metadata(title: str = None, raw: bool = False):
continue
scanned += 1
candidates.append(item)
log.debug(f'CivitAI search metadata: type={title if type(title) == str else "all"} workers={max_workers} skip={len(re_skip)} items={len(candidates)}')
log.debug(f'CivitAI search metadata: type={title if isinstance(title, str) else "all"} workers={max_workers} skip={len(re_skip)} items={len(candidates)}')
import concurrent
with concurrent.futures.ThreadPoolExecutor(max_workers=max_workers) as executor:
future_items = {}
@@ -283,5 +226,5 @@ def civit_search_metadata(title: str = None, raw: bool = False):
yield results if raw else create_search_metadata_table(results)
t1 = time.time()
log.debug(f'CivitAI search metadata: scanned={scanned} skipped={skipped} time={t1-t0:.2f}')
log.debug(f'CivitAI search metadata: scanned={scanned} skipped={skipped} time={t1 - t0:.2f}')
yield results if raw else create_search_metadata_table(results)
+159
View File
@@ -0,0 +1,159 @@
from pydantic import BaseModel, Field, validator # pylint: disable=no-name-in-module
class CivitImage(BaseModel):
class Config:
allow_population_by_field_name = True
id: int = 0
url: str = ""
width: int = 0
height: int = 0
type: str = "Unknown"
nsfw_level: int = Field(0, alias="nsfwLevel")
hash: str | None = None
meta: dict | None = None
class CivitFileHashes(BaseModel):
class Config:
allow_population_by_field_name = True
sha256: str | None = Field(None, alias="SHA256")
autov1: str | None = Field(None, alias="AutoV1")
autov2: str | None = Field(None, alias="AutoV2")
autov3: str | None = Field(None, alias="AutoV3")
crc32: str | None = Field(None, alias="CRC32")
blake3: str | None = Field(None, alias="BLAKE3")
class CivitFile(BaseModel):
class Config:
allow_population_by_field_name = True
id: int = 0
name: str = "Unknown"
type: str = "Unknown"
size_kb: float = Field(0, alias="sizeKB")
hashes: CivitFileHashes = Field(default_factory=CivitFileHashes)
download_url: str = Field("", alias="downloadUrl")
primary: bool | None = None
class CivitStats(BaseModel):
class Config:
allow_population_by_field_name = True
download_count: int = Field(0, alias="downloadCount")
favorite_count: int = Field(0, alias="favoriteCount")
thumb_up_count: int = Field(0, alias="thumbsUpCount")
thumb_down_count: int = Field(0, alias="thumbsDownCount")
comment_count: int = Field(0, alias="commentCount")
rating_count: int = Field(0, alias="ratingCount")
rating: float = 0
class CivitVersion(BaseModel):
class Config:
allow_population_by_field_name = True
id: int = 0
model_id: int = Field(0, alias="modelId")
name: str = "Unknown"
base_model: str = Field("Unknown", alias="baseModel")
published_at: str | None = Field(None, alias="publishedAt")
availability: str = "Unknown"
description: str | None = None
trained_words: list[str] = Field(default_factory=list, alias="trainedWords")
stats: CivitStats = Field(default_factory=CivitStats)
files: list[CivitFile] = Field(default_factory=list)
images: list[CivitImage] = Field(default_factory=list)
nsfw_level: int = Field(0, alias="nsfwLevel")
download_url: str = Field("", alias="downloadUrl")
class CivitCreator(BaseModel):
class Config:
allow_population_by_field_name = True
username: str = "Unknown"
image: str | None = None
class CivitModel(BaseModel):
class Config:
allow_population_by_field_name = True
id: int = 0
type: str = "Unknown"
name: str = "Unknown"
description: str | None = None
tags: list[str] = Field(default_factory=list)
nsfw: bool = False
nsfw_level: int = Field(0, alias="nsfwLevel")
availability: str = "Unknown"
stats: CivitStats = Field(default_factory=CivitStats)
creator: CivitCreator = Field(default_factory=CivitCreator)
versions: list[CivitVersion] = Field(default_factory=list, alias="modelVersions")
allow_no_credit: bool = Field(True, alias="allowNoCredit")
allow_commercial_use: list[str] = Field(default_factory=list, alias="allowCommercialUse")
allow_derivatives: bool = Field(True, alias="allowDerivatives")
allow_different_license: bool = Field(True, alias="allowDifferentLicense")
@validator('allow_commercial_use', pre=True)
def coerce_commercial_use(cls, v): # pylint: disable=no-self-argument
if isinstance(v, str):
return [v] if v else []
return v
class CivitSearchMetadata(BaseModel):
class Config:
allow_population_by_field_name = True
next_page: str | None = Field(None, alias="nextPage")
current_page: int | None = Field(None, alias="currentPage")
page_size: int | None = Field(None, alias="pageSize")
total_pages: int | None = Field(None, alias="totalPages")
total_items: int | None = Field(None, alias="totalItems")
next_cursor: str | None = Field(None, alias="nextCursor")
class CivitSearchResponse(BaseModel):
class Config:
allow_population_by_field_name = True
items: list[CivitModel] = Field(default_factory=list)
metadata: CivitSearchMetadata = Field(default_factory=CivitSearchMetadata)
request_url: str | None = Field(None, alias="requestUrl")
class CivitTag(BaseModel):
class Config:
allow_population_by_field_name = True
name: str = ""
model_count: int = Field(0, alias="modelCount")
link: str = ""
class CivitTagResponse(BaseModel):
class Config:
allow_population_by_field_name = True
items: list[CivitTag] = Field(default_factory=list)
metadata: CivitSearchMetadata = Field(default_factory=CivitSearchMetadata)
class CivitCreatorItem(BaseModel):
class Config:
allow_population_by_field_name = True
username: str = ""
model_count: int = Field(0, alias="modelCount")
link: str = ""
image: str | None = None
class CivitCreatorResponse(BaseModel):
class Config:
allow_population_by_field_name = True
items: list[CivitCreatorItem] = Field(default_factory=list)
metadata: CivitSearchMetadata = Field(default_factory=CivitSearchMetadata)
class CivitUserProfile(BaseModel):
class Config:
allow_population_by_field_name = True
id: int = 0
username: str = ""
image: str | None = None
profile_picture: str | None = Field(None, alias="profilePicture")
+46 -172
View File
@@ -1,181 +1,70 @@
from dataclasses import dataclass
import os
import json
import time
from installer import install
from modules.logger import log
from installer import log
from modules.civitai.client_civitai import client
from modules.civitai.models_civitai import CivitModel, CivitSearchResponse
full_dct = False
full_html = False
# Hardcoded fallback list — used by Gradio UI if discover_options() fails
base_models = ['', 'AuraFlow', 'Chroma', 'CogVideoX', 'Flux.1 S', 'Flux.1 D', 'Flux.1 Krea', 'Flux.1 Kontext', 'Flux.2 D', 'HiDream', 'Hunyuan 1', 'Hunyuan Video', 'Illustrious', 'Kolors', 'LTXV', 'Lumina', 'Mochi', 'NoobAI', 'PixArt a', 'PixArt E', 'Pony', 'Pony V7', 'Qwen', 'SD 1.4', 'SD 1.5', 'SD 1.5 LCM', 'SD 1.5 Hyper', 'SD 2.0', 'SD 2.1', 'SDXL 1.0', 'SDXL Lightning', 'SDXL Hyper', 'Wan Video 1.3B t2v', 'Wan Video 14B t2v', 'Wan Video 14B i2v 480p', 'Wan Video 14B i2v 720p', 'Wan Video 2.2 TI2V-5B', 'Wan Video 2.2 I2V-A14B', 'Wan Video 2.2 T2V-A14B', 'Wan Video 2.5 T2V', 'Wan Video 2.5 I2V', 'ZImageTurbo', 'Other']
@dataclass
class ModelImage:
def __init__(self, dct: dict):
if isinstance(dct, str):
dct = json.loads(dct)
self.id: int = dct.get('id', 0)
self.url: str = dct.get('url', '')
self.width: int = dct.get('width', 0)
self.height: int = dct.get('height', 0)
self.type: str = dct.get('type', 'Unknown')
self.dct: dict = dct if full_dct else {}
def __str__(self):
return f'ModelImage(id={self.id} url="{self.url}" width={self.width} height={self.height} type="{self.type}")'
@dataclass
class ModelFile:
def __init__(self, dct: dict):
if isinstance(dct, str):
dct = json.loads(dct)
self.id: int = dct.get('id', 0)
self.size: int = int(1024 * dct.get('sizeKB', 0))
self.name: str = dct.get('name', 'Unknown')
self.type: str = dct.get('type', 'Unknown')
self.hashes: list[str] = [str(h) for h in dct.get('hashes', {}).values()]
self.url: str = dct.get('downloadUrl', '')
self.dct: dict = dct if full_dct else {}
def __str__(self):
return f'ModelFile(id={self.id} name="{self.name}" size={self.size} type="{self.type}" url="{self.url}")'
@dataclass
class ModelVersion:
def __init__(self, dct: dict):
import bs4
if isinstance(dct, str):
dct = json.loads(dct)
self.id: int = dct.get('id', 0)
self.name: str = dct.get('name', 'Unknown')
self.base: str = dct.get('baseModel', 'Unknown')
self.mtime: str = dct.get('publishedAt', '')
self.downloads: int = dct.get('stats', {}).get('downloadCount', 0)
self.availability: str = dct.get('availability', 'Unknown')
self.html: str = dct.get('description', '') or '' if full_html else ''
self.desc: str = bs4.BeautifulSoup(dct.get('description', '') or '', features="html.parser").get_text()
self.files = [ModelFile(f) for f in dct.get('files', [])]
self.images = [ModelImage(i) for i in dct.get('images', [])]
self.dct: dict = dct if full_dct else {}
def __str__(self):
return f'ModelVersion(id={self.id} name="{self.name}" base="{self.base}" mtime="{self.mtime}" downloads={self.downloads} availability={self.availability} desc="{self.desc[:30]}...")'
@dataclass
class Model:
def __init__(self, dct: dict):
import bs4
if isinstance(dct, str):
dct = json.loads(dct)
self.id: int = dct.get('id', 0)
self.url: str = f'https://civitai.com/models/{self.id}'
self.type: str = dct.get('type', 'Unknown')
self.name: str = dct.get('name', 'Unknown')
self.html: str = dct.get('description', '') or '' if full_html else ''
self.desc: str = bs4.BeautifulSoup(dct.get('description', '') or '', features="html.parser").get_text()
self.tags: list[str] = dct.get('tags', [])
self.nsfw: bool = dct.get('nsfw', False)
self.level: str = dct.get('nsfwLevel', 0)
self.availability: str = dct.get('availability', 'Unknown')
self.downloads: int = dct.get('stats', {}).get('downloadCount', 0)
self.creator: str = dct.get('creator', {}).get('username', 'Unknown')
self.versions: list[ModelVersion] = [ModelVersion(v) for v in dct.get('modelVersions', [])]
self.dct: dict = dct if full_dct else {}
def __str__(self):
return f'Model(id={self.id} type={self.type} name="{self.name}" versions={len(self.versions)} nsfw={self.nsfw}/{self.level} downloads={self.downloads} author="{self.creator}" tags={self.tags} desc="{self.desc[:30]}...")'
models: list[Model] = [] # global cache for civitai search results
def search_civitai(
query:str,
tag:str = '', # optional:tag name
types:str = '', # (Checkpoint, TextualInversion, Hypernetwork, AestheticGradient, LORA, Controlnet, Poses)
sort:str = '', # (Highest Rated, Most Downloaded, Newest)
period:str = '', # (AllTime, Year, Month, Week, Day)
nsfw:bool = None, # optional:bool
limit:int = 0,
base:str = '', # list
token:str = None,
exact:bool = True,
):
global models # pylint: disable=global-statement
import requests
from urllib.parse import urlencode
install('beautifulsoup4')
if len(query) == 0:
query: str,
tag: str = '',
types: str = '',
sort: str = '',
period: str = '',
nsfw: bool = None,
limit: int = 0,
base: str = '',
token: str = None,
exact: bool = True,
) -> list[CivitModel]:
if not query:
log.error('CivitAI: empty query')
return []
t0 = time.time()
dct = { 'query': query }
if len(tag) > 0:
dct['tag'] = tag
if nsfw is not None:
dct['nsfw'] = 'true' if nsfw else 'false'
if limit > 0:
dct['limit'] = limit
if len(types) > 0:
dct['types'] = types
if len(sort) > 0:
dct['sort'] = sort
if len(period) > 0:
dct['period'] = period
if len(base) > 0:
dct['baseModels'] = base
encoded = urlencode(dct)
headers = {}
if token is None:
token = os.environ.get('CIVITAI_TOKEN', None)
if token is not None and len(token) > 0:
headers['Authorization'] = f'Bearer {token}'
url = 'https://civitai.com/api/v1/models'
# Numeric query → single model fetch
if query.isnumeric():
uri = f'{url}/{query}'
else:
uri = f'{url}?{encoded}'
log.info(f'CivitAI request: uri="{uri}" dct={dct} token={token is not None}')
result = requests.get(uri, headers=headers, timeout=60)
if result.status_code != 200:
log.error(f'CivitAI: code={result.status_code} reason={result.reason} uri={result.url}')
model = client.get_model(int(query), token=token)
if model:
t1 = time.time()
log.info(f'CivitAI result: id={query} time={t1 - t0:.2f}')
return [model]
return []
all_models: list[Model] = []
exact_models: list[Model] = []
dct = result.json()
if 'items' not in dct:
items = [dct] # single model
else:
items = dct.get('items', [])
for item in items:
all_models.append(Model(item))
response: CivitSearchResponse = client.search_models(
query=query,
tag=tag,
types=types,
sort=sort,
period=period,
base_models=[base] if base else None,
nsfw=nsfw,
limit=limit if limit > 0 else 20,
token=token,
)
all_models = response.items
exact_models: list[CivitModel] = []
if exact:
q_lower = query.lower()
for model in all_models:
model_names = [model.name.lower()]
version_names = [v.name.lower() for v in model.versions]
file_names = [f.name.lower() for v in model.versions for f in v.files]
if any([query.lower() in name for name in model_names + version_names + file_names]): # noqa: C419 # pylint: disable=use-a-generator
names = [model.name.lower()]
names.extend(v.name.lower() for v in model.versions)
names.extend(f.name.lower() for v in model.versions for f in v.files)
if any(q_lower in name for name in names):
exact_models.append(model)
result = exact_models if exact_models else all_models
t1 = time.time()
log.info(f'CivitAI result: code={result.status_code} exact={len(exact_models)} total={len(models)} time={t1-t0:.2f}')
models = exact_models if len(exact_models) > 0 else all_models
return models
log.info(f'CivitAI result: exact={len(exact_models)} total={len(all_models)} time={t1 - t0:.2f}')
return result
def create_model_cards(all_models: list[Model]) -> str:
def create_model_cards(all_models: list[CivitModel]) -> str:
details = """
<div id="model-details">
</div>
@@ -197,25 +86,10 @@ def create_model_cards(all_models: list[Model]) -> str:
previews = []
for version in model.versions:
for image in version.images:
if image.url and len(image.url) > 0 and not image.url.lower().endswith('.mp4'):
if image.url and not image.url.lower().endswith('.mp4'):
previews.append(image.url)
if len(previews) == 0:
if not previews:
previews = ['/sdapi/v1/network/thumb?filename=html/missing.png']
all_cards += card.format(id=model.id, name=model.name, type=model.type, preview=previews[0])
html = details + cards.format(cards=all_cards)
return html
def print_models(all_models: list[Model]):
for model in all_models:
log.info(f' {model}')
log.trace('Model', model.dct)
for version in model.versions:
log.info(f' {version}')
log.trace('ModelVersion', version.dct)
for file in version.files:
log.info(f' {file}')
log.trace('ModelFile', file.dct)
for image in version.images:
log.info(f' {image}')
log.trace('ModelImage', image.dct)
+128
View File
@@ -0,0 +1,128 @@
import os
import json
import threading
from datetime import datetime
from modules.logger import log
def data_dir() -> str:
from modules import paths
return paths.data_path or paths.script_path
class UserList:
def __init__(self, filename: str):
self._filename = filename
self._items: list[str] = []
self._lock = threading.Lock()
self._load()
def _path(self) -> str:
return os.path.join(data_dir(), self._filename)
def _load(self):
path = self._path()
if os.path.isfile(path):
try:
with open(path, encoding='utf-8') as f:
data = json.load(f)
if isinstance(data, list):
self._items = data
except Exception as e:
log.warning(f'CivitAI userdata load error: file={path} {e}')
def _save(self):
path = self._path()
try:
os.makedirs(os.path.dirname(path), exist_ok=True)
with open(path, 'w', encoding='utf-8') as f:
json.dump(self._items, f, indent=2)
except Exception as e:
log.error(f'CivitAI userdata save error: file={path} {e}')
def add(self, name: str) -> bool:
with self._lock:
if name not in self._items:
self._items.append(name)
self._save()
return True
return False
def remove(self, name: str) -> bool:
with self._lock:
if name in self._items:
self._items.remove(name)
self._save()
return True
return False
def list(self) -> list[str]:
with self._lock:
return list(self._items)
def contains(self, name: str) -> bool:
with self._lock:
return name in self._items
class SearchHistory:
def __init__(self, filename: str, max_entries: int = 30):
self._filename = filename
self._max_entries = max_entries
self._entries: list[dict] = []
self._lock = threading.Lock()
self._load()
def _path(self) -> str:
return os.path.join(data_dir(), self._filename)
def _load(self):
path = self._path()
if os.path.isfile(path):
try:
with open(path, encoding='utf-8') as f:
data = json.load(f)
if isinstance(data, list):
self._entries = data
except Exception as e:
log.warning(f'CivitAI search history load error: file={path} {e}')
def _save(self):
path = self._path()
try:
os.makedirs(os.path.dirname(path), exist_ok=True)
with open(path, 'w', encoding='utf-8') as f:
json.dump(self._entries, f, indent=2)
except Exception as e:
log.error(f'CivitAI search history save error: file={path} {e}')
def add(self, search_type: str, term: str):
with self._lock:
entry = {
"type": search_type,
"term": term,
"timestamp": datetime.now().isoformat(),
}
# Remove duplicate if same type+term exists
self._entries = [e for e in self._entries if not (e.get('type') == search_type and e.get('term') == term)]
self._entries.insert(0, entry)
# Trim to max
if len(self._entries) > self._max_entries:
self._entries = self._entries[:self._max_entries]
self._save()
def list(self, search_type: str = None) -> list[dict]:
with self._lock:
if search_type:
return [e for e in self._entries if e.get('type') == search_type]
return list(self._entries)
def clear(self):
with self._lock:
self._entries.clear()
self._save()
bookmarks = UserList("civitai_bookmarks.json")
banned = UserList("civitai_banned.json")
search_history = SearchHistory("civitai_search_history.json")
+2 -2
View File
@@ -279,8 +279,8 @@ def get_closest_checkpoint_match(s: str) -> CheckpointInfo:
# civitai search
if shared.opts.sd_checkpoint_autodownload and s.startswith("https://civitai.com/api/download/models"):
from modules.civitai.download_civitai import download_civit_model_thread
fn = download_civit_model_thread(model_name=None, model_url=s, model_path='', model_type='Model', token=None)
from modules.civitai.download_civitai import download_civit_model
fn = download_civit_model(model_url=s, model_name='', model_path='', model_type='Model', token=shared.opts.civitai_token)
if fn is not None:
checkpoint_info = CheckpointInfo(fn)
log.debug(f'Search model: name="{s}" matched="{checkpoint_info.path}" type=civitai')