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