diff --git a/modules/civitai/api_civitai.py b/modules/civitai/api_civitai.py index 4041f3f36..520a5ce43 100644 --- a/modules/civitai/api_civitai.py +++ b/modules/civitai/api_civitai.py @@ -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') diff --git a/modules/civitai/client_civitai.py b/modules/civitai/client_civitai.py new file mode 100644 index 000000000..ce6f84625 --- /dev/null +++ b/modules/civitai/client_civitai.py @@ -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() diff --git a/modules/civitai/download_civitai.py b/modules/civitai/download_civitai.py index 205d9c243..b2e809d06 100644 --- a/modules/civitai/download_civitai.py +++ b/modules/civitai/download_civitai.py @@ -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 diff --git a/modules/civitai/filemanage_civitai.py b/modules/civitai/filemanage_civitai.py new file mode 100644 index 000000000..febbcf361 --- /dev/null +++ b/modules/civitai/filemanage_civitai.py @@ -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 diff --git a/modules/civitai/metadata_civitai.py b/modules/civitai/metadata_civitai.py index d3594c3f7..57c4e082e 100644 --- a/modules/civitai/metadata_civitai.py +++ b/modules/civitai/metadata_civitai.py @@ -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 = """ @@ -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) diff --git a/modules/civitai/models_civitai.py b/modules/civitai/models_civitai.py new file mode 100644 index 000000000..2a1d74adc --- /dev/null +++ b/modules/civitai/models_civitai.py @@ -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") diff --git a/modules/civitai/search_civitai.py b/modules/civitai/search_civitai.py index ee2796b86..07fc7f809 100644 --- a/modules/civitai/search_civitai.py +++ b/modules/civitai/search_civitai.py @@ -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 = """
@@ -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) diff --git a/modules/civitai/userdata_civitai.py b/modules/civitai/userdata_civitai.py new file mode 100644 index 000000000..d990699a2 --- /dev/null +++ b/modules/civitai/userdata_civitai.py @@ -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") diff --git a/modules/sd_checkpoint.py b/modules/sd_checkpoint.py index 20da337fe..f6cd40fa3 100644 --- a/modules/sd_checkpoint.py +++ b/modules/sd_checkpoint.py @@ -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')