| import os |
| import random |
| import requests |
| import hashlib |
| import re |
| from typing import Sequence, Mapping, Any, Union, Set |
| from pathlib import Path |
| import shutil |
|
|
| import gradio as gr |
| from huggingface_hub import hf_hub_download, constants as hf_constants |
| import torch |
| import numpy as np |
| from PIL import Image, ImageChops |
| import yaml |
|
|
| from core.settings import * |
|
|
| MODELS_ROOT_DIR = "ComfyUI/models" |
|
|
|
|
| class UniqueKeyLoader(yaml.SafeLoader): |
| """ |
| A custom YAML loader that handles duplicate keys by grouping their values into a list. |
| """ |
| def construct_mapping(self, node, deep=False): |
| mapping = [] |
| for key_node, value_node in node.value: |
| key = self.construct_object(key_node, deep=deep) |
| value = self.construct_object(value_node, deep=deep) |
| mapping.append((key, value)) |
| |
| result = {} |
| for k, v in mapping: |
| if k in result: |
| if isinstance(result[k], list): |
| result[k].append(v) |
| else: |
| result[k] = [result[k], v] |
| else: |
| result[k] = v |
| return result |
|
|
| UniqueKeyLoader.add_constructor(yaml.resolver.BaseResolver.DEFAULT_MAPPING_TAG, UniqueKeyLoader.construct_mapping) |
|
|
| def save_uploaded_file_with_hash(file_obj: gr.File, target_dir: str) -> str: |
| if not file_obj: |
| return "" |
| |
| temp_path = file_obj.name |
| |
| sha256 = hashlib.sha256() |
| with open(temp_path, 'rb') as f: |
| for block in iter(lambda: f.read(65536), b''): |
| sha256.update(block) |
| |
| file_hash = sha256.hexdigest() |
| _, extension = os.path.splitext(temp_path) |
| hashed_filename = f"{file_hash}{extension.lower()}" |
| |
| dest_path = os.path.join(target_dir, hashed_filename) |
| |
| os.makedirs(target_dir, exist_ok=True) |
| if not os.path.exists(dest_path): |
| shutil.copy(temp_path, dest_path) |
| print(f"✅ Saved uploaded file as: {dest_path}") |
| else: |
| print(f"ℹ️ File already exists (deduplicated): {dest_path}") |
| |
| return hashed_filename |
|
|
| def bytes_to_gb(byte_size: int) -> float: |
| if byte_size is None or byte_size == 0: |
| return 0.0 |
| return round(byte_size / (1024 ** 3), 2) |
|
|
| def get_directory_size(path: str) -> int: |
| total_size = 0 |
| if not os.path.exists(path): |
| return 0 |
| try: |
| for dirpath, _, filenames in os.walk(path): |
| for f in filenames: |
| fp = os.path.join(dirpath, f) |
| if os.path.isfile(fp) and not os.path.islink(fp): |
| total_size += os.path.getsize(fp) |
| except OSError as e: |
| print(f"Warning: Could not access {path} to calculate size: {e}") |
| return total_size |
|
|
| def get_value_at_index(obj: Union[Sequence, Mapping], index: int) -> Any: |
| try: |
| return obj[index] |
| except (KeyError, IndexError): |
| try: |
| return obj["result"][index] |
| except (KeyError, IndexError): |
| return None |
|
|
| def sanitize_prompt(prompt: str) -> str: |
| if not isinstance(prompt, str): |
| return "" |
| return "".join(char for char in prompt if char.isprintable() or char in ('\n', '\t')) |
|
|
| def sanitize_id(input_id: str) -> str: |
| if not isinstance(input_id, str): |
| return "" |
| input_id = input_id.strip() |
| if "civitai" in input_id.lower(): |
| version_match = re.search(r'modelVersionId=(\d+)', input_id) |
| if version_match: |
| return version_match.group(1) |
| model_match = re.search(r'/models/(\d+)', input_id) |
| if model_match: |
| return model_match.group(1) |
| return re.sub(r'[^0-9]', '', input_id) |
|
|
| def sanitize_url(url: str) -> str: |
| if not isinstance(url, str): |
| raise ValueError("URL must be a string.") |
| url = url.strip() |
| if not re.match(r'^https?://[^\s/$.?#].[^\s]*$', url): |
| raise ValueError("Invalid URL format or scheme. Only HTTP and HTTPS are allowed.") |
| return url |
|
|
| def sanitize_filename(filename: str) -> str: |
| if not isinstance(filename, str): |
| return "" |
| sanitized = filename.replace('..', '') |
| sanitized = re.sub(r'[^\w\.\-]', '_', sanitized) |
| return sanitized.lstrip('/\\') |
|
|
| def get_civitai_file_info(version_id: str) -> dict | None: |
| api_url = f"https://civitai.com/api/v1/model-versions/{version_id}" |
| try: |
| response = requests.get(api_url, timeout=10) |
| response.raise_for_status() |
| data = response.json() |
| |
| model_type = data.get('model', {}).get('type') |
| |
| result_file = None |
| for file_data in data.get('files', []): |
| if file_data.get('type') == 'Model' and file_data['name'].endswith(('.safetensors', '.pt', '.bin')): |
| result_file = file_data.copy() |
| break |
| |
| if not result_file and data.get('files'): |
| result_file = data['files'][0].copy() |
| |
| if result_file: |
| result_file['model_type'] = model_type |
| return result_file |
| except Exception: |
| return None |
|
|
| def download_file(url: str, save_path: str, api_key: str = None, progress=None, desc: str = "") -> str: |
| if os.path.exists(save_path): |
| return f"File already exists: {os.path.basename(save_path)}" |
| |
| headers = {'Authorization': f'Bearer {api_key}'} if api_key and api_key.strip() else {} |
| try: |
| if progress: |
| progress(0, desc=desc) |
| |
| response = requests.get(url, stream=True, headers=headers, timeout=15) |
| response.raise_for_status() |
| total_size = int(response.headers.get('content-length', 0)) |
| |
| with open(save_path, "wb") as f: |
| downloaded = 0 |
| for chunk in response.iter_content(chunk_size=8192): |
| f.write(chunk) |
| if progress and total_size > 0: |
| downloaded += len(chunk) |
| progress(downloaded / total_size, desc=desc) |
| return f"Successfully downloaded: {os.path.basename(save_path)}" |
| except Exception as e: |
| if os.path.exists(save_path): |
| os.remove(save_path) |
| return f"Download failed for {os.path.basename(save_path)}: {e}" |
|
|
| def get_lora_path(source: str, id_or_url: str, civitai_key: str, progress) -> tuple[str | None, str]: |
| if not id_or_url or not id_or_url.strip(): |
| return None, "No ID/URL provided." |
|
|
| try: |
| if source == "Civitai": |
| version_id = sanitize_id(id_or_url) |
| if not version_id: |
| return None, "Invalid Civitai ID provided. Must be numeric." |
| |
| file_info = get_civitai_file_info(version_id) |
| if file_info: |
| model_type = file_info.get('model_type') |
| if model_type and model_type.lower() == 'checkpoint': |
| return None, f"Invalid Civitai model type '{model_type}' for LoRA. Checkpoint models are not allowed." |
| |
| filename = sanitize_filename(f"civitai_{version_id}.safetensors") |
| local_path = os.path.join(LORA_DIR, filename) |
| api_key_to_use = civitai_key |
| source_name = f"Civitai ID {version_id}" |
| elif source == "Hugging Face": |
| parts = id_or_url.strip().split('/') |
| if len(parts) < 3: |
| return None, "Invalid Hugging Face path. Format: repo_owner/repo_name/filename" |
| repo_id = f"{parts[0]}/{parts[1]}" |
| repo_file_path = "/".join(parts[2:]) |
| unique_name = id_or_url.strip().replace('/', '_') |
| filename = sanitize_filename(unique_name) |
| local_path = os.path.join(LORA_DIR, filename) |
| source_name = f"HF {repo_file_path}" |
| else: |
| return None, "Invalid source." |
|
|
| except ValueError as e: |
| return None, f"Input validation failed: {e}" |
|
|
| if os.path.lexists(local_path): |
| if not os.path.exists(local_path): |
| os.remove(local_path) |
| else: |
| return local_path, "File already exists." |
|
|
| if source == "Civitai": |
| if not file_info or not file_info.get('downloadUrl'): |
| return None, f"Could not get download link for {source_name}." |
|
|
| status = download_file(file_info['downloadUrl'], local_path, api_key_to_use, progress=progress, desc=f"Downloading {source_name}") |
| return (local_path, status) if "Successfully" in status else (None, status) |
| elif source == "Hugging Face": |
| try: |
| if progress and callable(progress): progress(0, desc=f"Downloading {source_name}") |
| cached_path = hf_hub_download(repo_id=repo_id, filename=repo_file_path, token=os.environ.get("HF_TOKEN")) |
| os.makedirs(LORA_DIR, exist_ok=True) |
| if os.path.lexists(local_path): |
| if not os.path.exists(local_path): |
| try: |
| os.remove(local_path) |
| except OSError: |
| pass |
| if not os.path.exists(local_path): |
| try: |
| os.symlink(cached_path, local_path) |
| except (OSError, NotImplementedError): |
| shutil.copyfile(cached_path, local_path) |
| if progress and callable(progress): progress(1.0, desc=f"Downloaded {source_name}") |
| return local_path, f"Successfully downloaded: {filename}" |
| except Exception as e: |
| return None, f"Hugging Face download failed: {e}" |
|
|
|
|
| def _ensure_model_downloaded(display_name: str, progress=gr.Progress()): |
| if display_name not in ALL_MODEL_MAP: |
| for cat_dir in CATEGORY_TO_DIR_MAP.values(): |
| check_path = os.path.join(cat_dir, display_name) |
| if os.path.exists(check_path): |
| return display_name |
| raise ValueError(f"Model '{display_name}' not found in configuration.") |
|
|
| model_info = ALL_MODEL_MAP[display_name] |
| repo_filename = model_info[1] |
| base_filename = os.path.basename(repo_filename) |
|
|
| download_info = ALL_FILE_DOWNLOAD_MAP.get(base_filename) |
| if not download_info: |
| raise gr.Error(f"Model '{base_filename}' not found in file_list.yaml. Cannot download.") |
|
|
| category = download_info.get("category") |
| dest_dir = CATEGORY_TO_DIR_MAP.get(category) |
| |
| if not dest_dir: |
| raise ValueError(f"Unknown YAML category '{category}' for '{base_filename}'.") |
| |
| dest_path = os.path.join(dest_dir, base_filename) |
|
|
| if os.path.lexists(dest_path): |
| if not os.path.exists(dest_path): |
| print(f"⚠️ Found and removed broken symlink: {dest_path}") |
| os.remove(dest_path) |
| else: |
| return base_filename |
|
|
| source = download_info.get("source") |
| try: |
| progress(0, desc=f"Downloading: {base_filename}") |
| |
| if source == "hf": |
| repo_id = download_info.get("repo_id") |
| hf_filename = download_info.get("repository_file_path", base_filename) |
| if not repo_id: |
| raise ValueError(f"repo_id is missing for HF model '{base_filename}'") |
| |
| cached_path = hf_hub_download(repo_id=repo_id, filename=hf_filename, token=os.environ.get("HF_TOKEN")) |
| os.makedirs(dest_dir, exist_ok=True) |
| os.symlink(cached_path, dest_path) |
| print(f"✅ Symlinked '{cached_path}' to '{dest_path}'") |
|
|
| elif source == "civitai": |
| model_version_id = download_info.get("model_version_id") |
| if not model_version_id: |
| raise ValueError(f"model_version_id is missing for Civitai model '{base_filename}'") |
| |
| file_info = get_civitai_file_info(model_version_id) |
| if not file_info or not file_info.get('downloadUrl'): |
| raise ConnectionError(f"Could not get download URL for Civitai model version ID {model_version_id}") |
| |
| status = download_file( |
| file_info['downloadUrl'], dest_path, api_key=os.environ.get("CIVITAI_API_KEY", ""), progress=progress, desc=f"Downloading: {base_filename}" |
| ) |
| if "Failed" in status: |
| raise ConnectionError(status) |
| else: |
| raise NotImplementedError(f"Download source '{source}' is not implemented for '{base_filename}'") |
| |
| progress(1.0, desc=f"Downloaded: {base_filename}") |
|
|
| except Exception as e: |
| if os.path.lexists(dest_path): |
| try: |
| os.remove(dest_path) |
| except OSError: pass |
| raise gr.Error(f"Failed to download and link '{display_name}': {e}") |
| |
| return base_filename |
|
|
|
|
| def ensure_file_downloaded(filename: str, progress=None): |
| if not filename or filename == "None": |
| return |
|
|
| download_info = ALL_FILE_DOWNLOAD_MAP.get(filename) |
| if not download_info: |
| print(f"⚠️ Warning: File '{filename}' not found in configuration (file_list.yaml). Cannot download.") |
| return |
|
|
| category = download_info.get("category", "loras") |
| dest_dir = CATEGORY_TO_DIR_MAP.get(category, LORA_DIR) |
| dest_path = os.path.join(dest_dir, filename) |
|
|
| if os.path.lexists(dest_path): |
| if not os.path.exists(dest_path): |
| print(f"⚠️ Found and removed broken symlink: {dest_path}") |
| os.remove(dest_path) |
| else: |
| return |
|
|
| source = download_info.get("source") |
| try: |
| if source == "hf": |
| repo_id = download_info.get("repo_id") |
| repo_filename = download_info.get("repository_file_path", filename) |
| if not repo_id: |
| raise ValueError("repo_id is missing for Hugging Face download.") |
|
|
| if progress and callable(progress): |
| progress(0, desc=f"Downloading: {filename}") |
| cached_path = hf_hub_download(repo_id=repo_id, filename=repo_filename, token=os.environ.get("HF_TOKEN")) |
| os.makedirs(dest_dir, exist_ok=True) |
| os.symlink(cached_path, dest_path) |
| print(f"✅ Symlinked '{cached_path}' to '{dest_path}'") |
| if progress and callable(progress): |
| progress(1.0, desc=f"Downloaded: {filename}") |
|
|
| elif source == "civitai": |
| model_version_id = download_info.get("model_version_id") |
| if not model_version_id: |
| raise ValueError("model_version_id is missing for Civitai download.") |
|
|
| file_info = get_civitai_file_info(model_version_id) |
| if not file_info or not file_info.get('downloadUrl'): |
| raise ConnectionError(f"Could not get download URL for Civitai model version ID {model_version_id}") |
|
|
| status = download_file( |
| file_info['downloadUrl'], |
| dest_path, |
| api_key=os.environ.get("CIVITAI_API_KEY", ""), |
| progress=progress, |
| desc=f"Downloading: {filename}" |
| ) |
| if "Failed" in status: |
| raise ConnectionError(status) |
| else: |
| raise NotImplementedError(f"Download source '{source}' is not implemented for '{filename}'.") |
|
|
| except Exception as e: |
| if os.path.lexists(dest_path): |
| try: |
| os.remove(dest_path) |
| except OSError: |
| pass |
| raise gr.Error(f"Failed to download file '{filename}': {e}") |
|
|
|
|
| def get_model_generation_defaults(model_display_name: str, model_type: str, defaults_config: dict): |
| final_defaults = { |
| 'steps': 25, 'cfg': 7.0, 'sampler_name': 'euler', 'scheduler': 'simple', |
| 'positive_prompt': '', 'negative_prompt': '' |
| } |
|
|
| if 'Default' in defaults_config: |
| final_defaults.update(defaults_config['Default']) |
|
|
| model_type_key = next((key for key in defaults_config if key.lower().replace(" ", "-").replace(".", "") == model_type.lower()), None) |
| if model_type_key: |
| model_type_config = defaults_config[model_type_key] |
| if '_defaults' in model_type_config: |
| final_defaults.update(model_type_config['_defaults']) |
| |
| if model_display_name in model_type_config: |
| final_defaults.update(model_type_config[model_display_name]) |
|
|
| return final_defaults |
|
|
| def get_filename_prefix() -> str: |
| import time |
| return f"LTX2.5_{int(time.time())}" |
|
|
| def get_media_metadata(file_obj, is_video=False): |
| default_video_meta = {'width': 0, 'height': 0, 'fps': 24, 'duration': 0} |
| default_image_meta = {'width': 0, 'height': 0, 'fps': 24} |
|
|
| if file_obj is None: |
| return default_video_meta if is_video else default_image_meta |
|
|
| if isinstance(file_obj, str) and os.path.exists(file_obj): |
| try: |
| import torchaudio |
| info = torchaudio.info(file_obj) |
| duration = info.num_frames / float(info.sample_rate) if info.sample_rate else 0 |
| if duration > 0: |
| return {'width': 0, 'height': 0, 'fps': 0, 'duration': duration} |
| except Exception: |
| pass |
|
|
| try: |
| import av |
| with av.open(file_obj) as container: |
| duration_sec = float(container.duration / av.time_base) if container.duration is not None else 0 |
| return {'width': 0, 'height': 0, 'fps': 24, 'duration': duration_sec} |
| except Exception: |
| pass |
|
|
| try: |
| import imageio.v2 as iio |
| with iio.get_reader(file_obj) as reader: |
| meta = reader.get_meta_data() |
| size = meta.get('size', meta.get('source_size', (0, 0))) |
| width, height = size |
| fps = meta.get('fps', 24) |
| duration = meta.get('duration', 0) |
| return {'width': width, 'height': height, 'fps': fps, 'duration': duration} |
| except Exception: |
| pass |
|
|
| if is_video: |
| return default_video_meta |
| else: |
| if isinstance(file_obj, Image.Image): |
| width, height = file_obj.size |
| return {'width': width, 'height': height, 'fps': 24} |
| return default_image_meta |
|
|
| def save_temp_image(img): |
| if not isinstance(img, Image.Image): |
| return None |
| _PROJECT_ROOT = os.path.dirname(os.path.dirname(os.path.abspath(__file__))) |
| input_dir = os.path.join(_PROJECT_ROOT, "input") |
| os.makedirs(input_dir, exist_ok=True) |
| filename = f"temp_image_{random.randint(10000, 99999)}.png" |
| filepath = os.path.join(input_dir, filename) |
| img.save(filepath, "PNG") |
| return os.path.basename(filepath) |
|
|
| def save_temp_audio(audio_path): |
| if not audio_path or not os.path.exists(audio_path): |
| return None |
| _PROJECT_ROOT = os.path.dirname(os.path.dirname(os.path.abspath(__file__))) |
| input_dir = os.path.join(_PROJECT_ROOT, "input") |
| os.makedirs(input_dir, exist_ok=True) |
| ext = os.path.splitext(audio_path)[1] or ".wav" |
| filename = f"temp_audio_{random.randint(10000, 99999)}{ext}" |
| save_path = os.path.join(input_dir, filename) |
| shutil.copy(audio_path, save_path) |
| return os.path.basename(filename) |
|
|
| def save_temp_video(video_path): |
| if not video_path or not os.path.exists(video_path): |
| return None |
| _PROJECT_ROOT = os.path.dirname(os.path.dirname(os.path.abspath(__file__))) |
| input_dir = os.path.join(_PROJECT_ROOT, "input") |
| os.makedirs(input_dir, exist_ok=True) |
| ext = os.path.splitext(video_path)[1] or ".mp4" |
| filename = f"temp_video_{random.randint(10000, 99999)}{ext}" |
| save_path = os.path.join(input_dir, filename) |
| shutil.copy(video_path, save_path) |
| return os.path.basename(filename) |
|
|
| def handle_seed(seed_value: int, max_val: int = 2**32 - 1) -> int: |
| if seed_value == -1 or seed_value is None: |
| return random.randint(0, max_val) |
| return int(seed_value) |
|
|
| def process_lora_inputs(ui_values: dict, prefix: str = "", progress=None) -> list: |
| active_loras_for_gpu = [] |
|
|
| |
| lora_sources = ui_values.get(f'lora_sources_{prefix}', []) if prefix else [] |
| lora_ids = ui_values.get(f'lora_ids_{prefix}', []) if prefix else [] |
| lora_scales = ui_values.get(f'lora_scales_{prefix}', []) if prefix else [] |
|
|
| if isinstance(lora_sources, list) and isinstance(lora_ids, list): |
| for source, val, scale in zip(lora_sources, lora_ids, lora_scales): |
| scale_val = float(scale) if scale is not None else 1.0 |
| if scale_val > 0 and val and str(val).strip(): |
| lora_id = str(val).strip() |
| lora_filename = None |
| if source == "File": |
| lora_filename = sanitize_filename(lora_id) |
| local_path = os.path.join(LORA_DIR, lora_filename) |
| if not os.path.exists(local_path): |
| raise gr.Error(f"Uploaded LoRA file '{lora_id}' no longer exists on server. Please re-upload it.") |
| elif source in ("Civitai", "Hugging Face"): |
| local_path, status = get_lora_path(source, lora_id, os.environ.get("CIVITAI_API_KEY", ""), progress) |
| if local_path: |
| lora_filename = os.path.basename(local_path) |
| else: |
| raise gr.Error(f"Failed to prepare LoRA {lora_id}: {status}") |
|
|
| if lora_filename: |
| active_loras_for_gpu.append({ |
| "lora_name": lora_filename, |
| "strength_model": scale_val, |
| "strength_clip": scale_val |
| }) |
|
|
| |
| lora_data = ui_values.get('lora_data', []) |
| if lora_data and not active_loras_for_gpu: |
| sources, ids, scales, files = lora_data[0::4], lora_data[1::4], lora_data[2::4], lora_data[3::4] |
| for source, lora_id, scale, _ in zip(sources, ids, scales, files): |
| scale_val = float(scale) if scale is not None else 1.0 |
| if scale_val > 0 and lora_id and str(lora_id).strip(): |
| lora_id_str = str(lora_id).strip() |
| lora_filename = None |
| if source == "File": |
| lora_filename = sanitize_filename(lora_id_str) |
| local_path = os.path.join(LORA_DIR, lora_filename) |
| if not os.path.exists(local_path): |
| raise gr.Error(f"Uploaded LoRA file '{lora_id_str}' no longer exists on server. Please re-upload it.") |
| elif source in ("Civitai", "Hugging Face"): |
| local_path, status = get_lora_path(source, lora_id_str, os.environ.get("CIVITAI_API_KEY", ""), progress) |
| if local_path: |
| lora_filename = os.path.basename(local_path) |
| else: |
| raise gr.Error(f"Failed to prepare LoRA {lora_id_str}: {status}") |
|
|
| if lora_filename: |
| active_loras_for_gpu.append({ |
| "lora_name": lora_filename, |
| "strength_model": scale_val, |
| "strength_clip": scale_val |
| }) |
|
|
| |
| raw_loras = ui_values.get('loras', []) |
| if raw_loras and not active_loras_for_gpu and isinstance(raw_loras, list): |
| for item in raw_loras: |
| if isinstance(item, dict): |
| if "lora_name" in item: |
| active_loras_for_gpu.append(item) |
| else: |
| src = item.get("source", "Hugging Face") |
| val = item.get("lora_value") or item.get("id_or_url") or item.get("lora_id") |
| scale = item.get("scale", 1.0) |
| scale_val = float(scale) if scale is not None else 1.0 |
| if scale_val > 0 and val and str(val).strip(): |
| lora_id_str = str(val).strip() |
| lora_filename = None |
| if src == "File": |
| lora_filename = sanitize_filename(lora_id_str) |
| local_path = os.path.join(LORA_DIR, lora_filename) |
| if not os.path.exists(local_path): |
| raise gr.Error(f"Uploaded LoRA file '{lora_id_str}' no longer exists on server. Please re-upload it.") |
| elif src in ("Civitai", "Hugging Face"): |
| local_path, status = get_lora_path(src, lora_id_str, os.environ.get("CIVITAI_API_KEY", ""), progress) |
| if local_path: |
| lora_filename = os.path.basename(local_path) |
| else: |
| raise gr.Error(f"Failed to prepare LoRA {lora_id_str}: {status}") |
| if lora_filename: |
| active_loras_for_gpu.append({ |
| "lora_name": lora_filename, |
| "strength_model": scale_val, |
| "strength_clip": scale_val |
| }) |
|
|
| return active_loras_for_gpu |