Files
arcum42-ComfyUI_SageUtils/utils/helpers.py
T

799 lines
28 KiB
Python

#Utility functions for use in the nodes.
import pathlib
import hashlib
import datetime
import torch
import logging
import folder_paths
import comfy.utils
from .model_cache import cache
from .helpers_civitai import *
from .constants import MODEL_FILE_EXTENSIONS
def str_to_bool(value):
if isinstance(value, bool):
return value
if isinstance(value, str):
value = value.lower()
if value in {'true', '1', 'yes'}:
return True
if value in {'false', '0', 'no'}:
return False
raise ValueError(f"Cannot convert {value} to boolean.")
# Not currently used.
def bool_to_str(value):
if isinstance(value, bool):
return "true" if value else "false"
if isinstance(value, str):
value = value.lower()
if value in {'true', '1', 'yes'}:
return "true"
if value in {'false', '0', 'no'}:
return "false"
raise ValueError(f"Cannot convert {value} to string representation of boolean.")
def name_from_path(path):
return pathlib.Path(path).name
def get_path_without_base(folder_type:str, path:str) -> str:
"""Get the base path for a given folder type and path."""
for base in folder_paths.get_folder_paths(folder_type):
if path.startswith(base):
path = path[len(base):].lstrip("/\\")
break
return path
def get_file_extension(path: str) -> str:
"""Get the file extension from a path."""
return path.split(".")[-1] if "." in path else ""
def has_model_extension(path: str) -> bool:
"""Check if a file path has a model extension."""
if not path:
return False
extension = '.' + get_file_extension(path).lower()
return extension in {ext.lower() for ext in MODEL_FILE_EXTENSIONS}
def is_model_file(path: str) -> bool:
"""Check if a file is a valid model file based on its extension."""
return has_model_extension(path)
# fast_hash_threshold=100*1024*1024 # 100MB threshold
def get_file_sha256(path, fast_hash_threshold=0): # Full hashing by default
"""
Calculate SHA256 hash of a file with optimizations for large files.
Args:
path: File path to hash
fast_hash_threshold: File size threshold (in bytes) above which to use fast hashing.
Set to 0 to always use full hashing.
Returns:
First 10 characters of the SHA256 hash
"""
print(f"Calculating hash for {path}")
import hashlib
import os
# Use a larger buffer size for better I/O performance
BUFFER_SIZE = 65536 # 64KB chunks
try:
file_size = os.path.getsize(path)
print(f"File size: {file_size / (1024*1024):.1f} MB")
m = hashlib.sha256()
# For very large files, use a sampling approach for speed
if fast_hash_threshold > 0 and file_size > fast_hash_threshold:
print(f"Large file detected, using fast hashing strategy")
# Sample strategy: hash beginning, middle, and end of file + file size + filename
# This creates a unique signature that's very unlikely to collide for different files
sample_size = min(BUFFER_SIZE * 4, file_size // 10) # Sample up to 256KB or 10% of file
with open(path, 'rb') as f:
# Hash file metadata first (size and name)
m.update(str(file_size).encode())
m.update(os.path.basename(path).encode())
# Hash beginning of file
beginning = f.read(sample_size)
m.update(beginning)
# Hash middle of file
if file_size > sample_size * 2:
middle_pos = file_size // 2 - sample_size // 2
f.seek(middle_pos)
middle = f.read(sample_size)
m.update(middle)
# Hash end of file
if file_size > sample_size * 3:
f.seek(-sample_size, 2) # Seek to sample_size bytes from end
end = f.read(sample_size)
m.update(end)
else:
# Full file hashing for smaller files or when requested
with open(path, 'rb') as f:
while True:
chunk = f.read(BUFFER_SIZE)
if not chunk:
break
m.update(chunk)
except (OSError, IOError) as e:
print(f"Error reading file {path}: {e}")
raise
full_hash = m.hexdigest()
print(f"Got full hash {full_hash}")
result = full_hash[:10]
print(f"Got hash {result}")
return result
def get_files_in_dir(input_dirs=None, extensions=None):
if extensions is None:
extensions = ".*"
if input_dirs is None:
raise ValueError("input_dirs cannot be None")
input_files = []
# Check if input_dirs is a tuple or a list. If not, make it a list.
if not isinstance(input_dirs, (list, tuple)):
input_dirs = [input_dirs]
for dir in input_dirs:
if dir is None or dir == "":
continue
if pathlib.Path(dir).exists():
file_list = pathlib.Path(dir).rglob("*")
for file in file_list:
if file.exists() and not file.is_dir():
file_path = ""
if file.suffix.lower() in extensions:
# Get path relative to the input directory
try:
file_path = str(file.relative_to(dir))
except ValueError:
file_path = str(file)
input_files.append(file_path)
input_files = sorted(set(input_files))
return input_files
def last_used(file_path):
cache.load()
if file_path in cache.hash:
last_used = cache.by_path(file_path).get("lastUsed", None)
if last_used is not None:
return datetime.datetime.fromisoformat(last_used)
else:
return None
else:
return None
def days_since_last_used(file_path):
was_last_used = last_used(file_path)
if was_last_used is not None:
now = datetime.datetime.now()
delta = now - was_last_used
return delta.days
else:
return 365
def get_file_modification_date(file_path):
try:
file_path = pathlib.Path(file_path)
if file_path.exists():
return datetime.datetime.fromtimestamp(file_path.stat().st_mtime)
else:
print(f"File {file_path} does not exist.")
return datetime.datetime.now()
except Exception as e:
print(f"Error getting modification date for {file_path}: {e}")
return datetime.datetime.now()
# Called *after* verifying the json pulled successfully.
# This function updates the cache with the information from the CivitAI json.
def update_cache_from_civitai_json(file_path, json_data, timestamp=True):
the_files = json_data.get("files", [])
hashes = {}
if len(the_files) > 0:
hashes = the_files[0].get("hashes", {})
update_available = True
latest_model = None
if json_data.get("modelId", None) is not None:
latest_model = get_latest_model_version(json_data["modelId"])
if latest_model == json_data["id"] or latest_model is None:
update_available = False
if latest_model is None:
latest_model = ""
file_cache = cache.by_path(file_path)
file_cache.update({
'civitai': "True",
'civitai_failed_count': 0,
'model': json_data.get("model", {}),
'name': json_data.get("name", ""),
'baseModel': json_data.get("baseModel", ""),
'id': json_data.get("id", ""),
'modelId': json_data.get("modelId", ""),
'update_available': update_available,
'update_version_id': latest_model,
'trainedWords': json_data.get("trainedWords", []),
'downloadUrl': json_data.get("downloadUrl", ""),
'hashes': hashes
})
if timestamp:
print("Updating timestamp.")
cache.update_last_used_by_path(file_path)
print("Successfully pulled metadata.")
def update_cache_without_civitai_json(file_path, hash, timestamp=True):
file_cache = cache.by_path(file_path)
print(f"Unable to find on civitai.")
file_cache['civitai'] = "False"
file_cache['civitai_failed_count'] = file_cache.get('civitai_failed_count', 0) + 1
file_cache['hash'] = hash
cache.update_last_used_by_path(file_path)
def add_file_to_cache(file_path, hash=None):
file_path = str(file_path)
print(f"Adding {file_path} to cache.")
if hash is None:
hash = get_file_sha256(file_path)
if file_path not in cache.hash:
cache.hash[file_path] = hash
if cache.info.get(hash, None) is None:
cache.info[hash] = {
'civitai': "False",
'update_available': False,
'update_version_id': "",
'hash': hash,
'lastUsed': datetime.datetime.now().isoformat()
}
print(f"Adding {file_path} to cache with hash {hash}.")
return hash
def recheck_hash(file_path, hash):
new_hash = get_file_sha256(file_path)
if new_hash != hash:
print(f"Hash mismatch. Using new hash.")
if file_path in cache.hash:
print(f"Updating cache for {file_path} with new hash {new_hash}.")
if new_hash not in cache.info:
if hash in cache.info:
cache.info[new_hash] = cache.info[hash]
else:
print(f"File {file_path} not in cache. Adding with new hash {new_hash}.")
add_file_to_cache(file_path, new_hash)
hash = new_hash
return hash
# Model path is inconsistent. Sometimes it is a list, and sometimes a string.
# Also, different models vary on whether they use an absolute or relative path.
# Loras and VAEs use a relative path from the loras folder. They are converted to absolute paths
# when calling this function. get_full_path_or_raise will raise if passed an absolute path.
# Unfortunately, this is set in the main program, not this node set.
def pull_and_update_model_timestamp(file_paths, model_type):
if not isinstance(file_paths, (list, tuple)):
file_paths = [file_paths]
#if model_type == "lora":
# file_paths = [folder_paths.get_full_path_or_raise("loras", lora) for lora in file_paths]
for path in file_paths:
pull_metadata(path, model_type=model_type)
update_model_timestamp(file_paths)
def update_model_timestamp(file_paths):
cache.load()
# If the file_paths isn't a list, make it a list.
if not isinstance(file_paths, (list, tuple)):
file_paths = [file_paths]
for path in file_paths:
if path in cache.hash:
cache.update_last_used_by_path(path)
cache.save()
def pull_metadata(file_paths, timestamp = True, force_all = False, pbar = None, model_type = None):
pull_json = True
metadata_days_recheck = 7
cache.load()
cache.backup_counter += 1
if cache.backup_counter >= cache.num_of_backups_to_keep:
cache.prune_all_backups()
cache.backup_counter = 0
if isinstance(file_paths, str):
file_paths = [file_paths]
if not file_paths:
logging.warning("No file paths provided.")
return
num_not_pulled = 0
for file_path in file_paths:
force = force_all
hash = cache.hash.get(str(file_path), None)
if hash is None:
logging.debug(f"Hash not found in cache for {file_path}. Adding to cache.")
hash = add_file_to_cache(file_path)
file_cache = cache.by_path(file_path)
last_used_date = datetime.datetime.fromisoformat(file_cache['lastUsed']) if 'lastUsed' in file_cache else None
modified = get_file_modification_date(file_path)
# If file was modified after last used, force metadata pull
if last_used_date is not None and modified is not None and modified > last_used_date:
logging.info(f"File was modified after last used. Pulling metadata.")
force = True
# Only skip pull if not forced and civitai is True and recently pulled
civitai_val = False
try:
civitai_val = str_to_bool(file_cache.get('civitai', False))
except:
civitai_val = False
if not force and civitai_val == True:
if days_since_last_used(file_path) <= metadata_days_recheck:
num_not_pulled += 1
pull_json = False
if file_cache.get('blacklist'):
if not force:
logging.warning(f"File {file_path} is blacklisted (previously not found). Skipping metadata pull.")
pull_json = False
# If force, recalculate hash before any API call
if force:
logging.debug(f"Force flag is set. Recalculating hash for {file_path}.")
hash = recheck_hash(file_path, hash)
if pull_json or force:
logging.debug(f"Currently pulling metadata for {file_path}.")
json = get_civitai_model_version_json_by_hash(hash)
if 'error' in json:
retried = False
dead_model = False
if 'civitai_error' in json:
if 'Model not found' in json['civitai_error'] or 'No model with id' in json['civitai_error']:
dead_model = True
# Try fallback with modelId if available
if dead_model is False:
if 'modelId' in file_cache:
logging.debug(f"Using cached model id {file_cache.get('id', None)}")
json = get_civitai_model_version_json_by_id(file_cache['id'])
retried = True
else:
logging.debug(f"No cached model id.")
if 'error' in json:
if retried:
logging.error(f"Error: {json['error']}")
if dead_model:
file_cache['blacklist'] = True
print(f"Unable to find on civitai.")
file_cache['civitai'] = "False"
file_cache['civitai_failed_count'] = file_cache.get('civitai_failed_count', 0) + 1
update_cache_without_civitai_json(file_path, hash, timestamp=timestamp)
if 'error' not in json:
update_cache_from_civitai_json(file_path, json, timestamp=timestamp)
else:
retries = file_cache.get('civitai_failed_count', 0)
retries += 1
file_cache['civitai_failed_count'] = retries
cache.hash[file_path] = hash
cache.info[hash] = file_cache
if model_type is not None:
file_cache['model_type'] = model_type
if pbar is not None:
pbar.update(1)
if num_not_pulled > 0:
logging.info(f"Metadata pull complete. Skipped {num_not_pulled} files checked within the last {metadata_days_recheck} days.")
cache.save()
def lora_to_string(lora_name, model_weight, clip_weight):
lora_string = ' <lora:' + str(pathlib.Path(lora_name).name) + ":" + str(model_weight) + ">" # + ":" + str(clip_weight)
return lora_string
def lora_to_prompt(lora_stack = None):
lora_info = ''
if lora_stack is None:
return ""
else:
if isinstance(lora_stack, tuple) or isinstance(lora_stack, list):
if not isinstance(lora_stack[0], (list, tuple)):
lora_stack = [lora_stack]
for lora in lora_stack:
lora_info += lora_to_string(lora[0], lora[1], lora[2])
return lora_info
def get_lora_hash(lora_name):
lora_path = folder_paths.get_full_path_or_raise("loras", lora_name)
pull_metadata(lora_path)
return cache.hash[lora_path]
def model_scan(the_path, force = False):
the_paths = the_path
print(f"the_paths: {the_paths}")
model_list = []
for dir in the_paths:
print(f"dir: {dir}")
result = list(p.resolve() for p in pathlib.Path(dir).glob("**/*") if p.suffix in MODEL_FILE_EXTENSIONS)
model_list.extend(result)
model_list = list(set(model_list))
model_list = [str(x) for x in model_list]
print(f"Scanning {len(model_list)} models for metadata.")
pbar = comfy.utils.ProgressBar(len(model_list))
pull_metadata(model_list, force_all=force, pbar=pbar)
def grab_model_list(model_type: str, extra_models: list[str] | None = None) -> list[str]:
"""Get a list of model names based on the model type, including extra models."""
model_list = folder_paths.get_filename_list(model_type)
if extra_models is None:
return model_list
for extra_model in extra_models:
extra = []
try:
extra = [x for x in folder_paths.get_filename_list(extra_model)]
except Exception as e:
extra = []
#print(f"Extra models for {extra_model}: {extra}")
model_list += extra
model_list = list(set(model_list))
model_list.sort() # Remove duplicates
#print(f"Final model list for {model_type}: {model_list}")
return model_list
# Module-level cache for get_model_list
_model_list_cache = {}
def get_model_list(model_type: str) -> list[str]:
"""Get a list of model names based on the model type, with in-memory cache (1 min)."""
import time
global _model_list_cache
now = time.time()
cache_entry = _model_list_cache.get(model_type)
if cache_entry:
cached_time, cached_list = cache_entry
if now - cached_time < 60:
return cached_list
# Not cached or cache expired, fetch fresh
if model_type == "checkpoints":
result = grab_model_list("checkpoints")
elif model_type == "unet":
result = grab_model_list("unet", ["diffusion_models", "unet_gguf"])
elif model_type == "vae":
result = grab_model_list("vae")
elif model_type == "clip":
result = grab_model_list("clip", ["text_encoders", "clip_gguf"])
elif model_type == "loras":
result = grab_model_list("loras")
else:
result = []
_model_list_cache[model_type] = (now, result)
return result
def normalize_prompt_weights(text):
"""
Parse and normalize text prompts with weight formatting.
Converts prompts with implicit weights (parentheses) and explicit weights (text):weight
into a normalized format using only explicit weight notation.
Weight rules:
- Default weight: 1.0 (no parentheses in output)
- Each layer of parentheses adds 0.1 to the weight
- Explicit format: (text):weight stacks with surrounding parentheses
- Weight 1.1 outputs as (text) without weight notation
- All other weights output as (text:weight)
Example input: "The (cat) (jumped over (the dog):1.5)"
Example output: "The (cat) (jumped over (the dog:1.5))"
Cleanup rules applied:
- Multiple spaces reduced to 1
- Commas always followed by a space
- No space after ( or before )
- Space between ) and (
- Max one empty line in a row
- Unmatched parentheses are automatically balanced
"""
import re
# Balance parentheses
open_count = text.count('(')
close_count = text.count(')')
if open_count > close_count:
text = text + ')' * (open_count - close_count)
elif close_count > open_count:
text = '(' * (close_count - open_count) + text
def parse_weighted_text(text):
"""Parse text and return list of (content, weight, is_explicit) tuples."""
result = []
i = 0
while i < len(text):
# Look for opening parenthesis
paren_start = text.find('(', i)
if paren_start == -1:
# No more parentheses, add remaining text
if i < len(text):
result.append((text[i:], 1.0, False))
break
# Add text before parenthesis
if paren_start > i:
result.append((text[i:paren_start], 1.0, False))
# Find matching closing parenthesis and parse content
paren_depth = 0
j = paren_start
while j < len(text):
if text[j] == '(':
paren_depth += 1
elif text[j] == ')':
paren_depth -= 1
if paren_depth == 0:
break
j += 1
if paren_depth != 0:
# Unmatched parenthesis, treat as regular text
result.append((text[i:], 1.0, False))
break
# Extract content between parentheses
content = text[paren_start + 1:j]
# Check for explicit weight specification
weight_match = re.match(r'^(.*?):\s*([0-9]*\.?[0-9]+)\s*$', content)
if weight_match:
# Explicit weight specified
inner_content = weight_match.group(1).strip()
weight = float(weight_match.group(2))
result.append((inner_content, weight, True))
else:
# Implicit weight: recurse to calculate based on nested parentheses
nested_parts = parse_weighted_text(content)
# Adjust weights: add 0.1 for this level, for both implicit and explicit weights
for nested_content, nested_weight, is_explicit in nested_parts:
# Add 0.1 to both implicit and explicit weights for this parenthesis level
result.append((nested_content, nested_weight + 0.1, is_explicit))
i = j + 1
return result
def format_output(parts):
"""Format parsed parts into normalized output."""
output = []
for content, weight, is_explicit in parts:
# Round weight to avoid floating-point precision issues
weight = round(weight, 10)
# Recursively handle nested content
if '(' in content:
formatted = format_output(parse_weighted_text(content))
if weight == 1.0:
output.append(formatted)
elif weight == 1.1:
output.append(f"({formatted})")
else:
output.append(f"({formatted}:{weight})")
else:
content = content.strip()
if content:
if weight == 1.0:
output.append(content)
elif weight == 1.1:
output.append(f"({content})")
else:
output.append(f"({content}:{weight})")
return ' '.join(output)
# Parse the text
parts = parse_weighted_text(text)
# Format output
result = format_output(parts)
# Apply text cleanup
# Multiple spaces to single space
result = re.sub(r' +', ' ', result)
# Commas always followed by space
result = re.sub(r',\s*', ', ', result)
# Clean up spacing around parentheses
result = re.sub(r'\(\s+', '(', result) # No space after (
result = re.sub(r'\s+\)', ')', result) # No space before )
result = re.sub(r'\)\s+', ')', result) # No space after ) initially
# Add space after ) if followed by alphanumeric or (
result = re.sub(r'\)(?=[a-zA-Z0-9(])', ') ', result)
# Handle line breaks - max one empty line
result = re.sub(r'\n\s*\n\s*\n+', '\n\n', result)
# Clean leading/trailing whitespace
result = result.strip()
return result
def clean_keywords(keywords):
keywords = set(filter(None, (x.strip() for x in keywords)))
return ', '.join(keywords)
def clean_text(text):
ret = normalize_prompt_weights(text)
return ret
def clean_if_needed(text, clean):
return clean_text(text) if clean and text is not None else text
def condition_text(clip, text = None):
zero_text = text is None
text = text or ""
tokens = clip.tokenize(text)
output = clip.encode_from_tokens(tokens, return_pooled=True, return_dict=True)
cond = output.pop("cond")
if zero_text:
pooled_output = output.get("pooled_output")
if pooled_output is not None:
output["pooled_output"] = torch.zeros_like(pooled_output)
return [[torch.zeros_like(cond), output]]
return [[cond, output]]
def get_save_file_path(filename_prefix: str = "text", filename_ext: str = "txt") -> str:
"""
Generate a safe file path for saving files with automatic counter increment.
Args:
filename_prefix: Base filename, can include date/time variables like %year%, %month%, etc.
filename_ext: File extension (without dot)
Returns:
Complete file path including directory and filename with counter
"""
def _extract_counter_from_filename(filename: str) -> tuple[int, str]:
"""Extract counter from existing filename to determine next counter value."""
base_name = pathlib.Path(filename_prefix).name
prefix_len = len(base_name)
if len(filename) <= prefix_len + 1:
return 0, filename[:prefix_len + 1]
prefix = filename[:prefix_len + 1]
try:
# Remove file extension first, then extract counter
filename_no_ext = pathlib.Path(filename).stem
counter_part = filename_no_ext[prefix_len + 1:]
digits = int(counter_part)
except (ValueError, IndexError):
digits = 0
return digits, prefix
def _replace_date_variables(text: str) -> str:
"""Replace date/time variables in the filename prefix."""
now = datetime.datetime.now()
replacements = {
"%year%": str(now.year),
"%month%": str(now.month).zfill(2),
"%day%": str(now.day).zfill(2),
"%hour%": str(now.hour).zfill(2),
"%minute%": str(now.minute).zfill(2),
"%second%": str(now.second).zfill(2),
}
for placeholder, value in replacements.items():
text = text.replace(placeholder, value)
return text
# Get output directory and process filename prefix
output_dir = folder_paths.get_output_directory()
if "%" in filename_prefix:
filename_prefix = _replace_date_variables(filename_prefix)
# Parse the filename prefix path
filename_prefix_path = pathlib.Path(filename_prefix)
subfolder = filename_prefix_path.parent
base_filename = filename_prefix_path.name
# Construct full output path
output_path = pathlib.Path(output_dir)
full_output_folder = output_path / subfolder
# Security check: ensure we're not saving outside the output directory
try:
full_output_folder.resolve().relative_to(output_path.resolve())
except ValueError:
error_msg = (
"ERROR: Saving outside the output folder is not allowed.\n"
f" Target folder: {full_output_folder.resolve()}\n"
f" Output directory: {output_path.resolve()}"
)
logging.error(error_msg)
raise ValueError(error_msg)
# Ensure output directory exists
full_output_folder.mkdir(parents=True, exist_ok=True)
# Find the next available counter
counter = 1
try:
existing_files = [f.name for f in full_output_folder.iterdir() if f.is_file()]
matching_counters = []
for file in existing_files:
digits, prefix = _extract_counter_from_filename(file)
# Check if this file matches our pattern (same base name and ends with underscore)
if (prefix[:-1].lower() == base_filename.lower() and
len(prefix) > 0 and prefix[-1] == "_"):
matching_counters.append(digits)
if matching_counters:
counter = max(matching_counters) + 1
except Exception as e:
logging.warning(f"Error finding existing files, using counter=1: {e}")
counter = 1
# Generate final filename
final_filename = f"{base_filename}_{counter:05d}.{filename_ext}"
return str(full_output_folder / final_filename)
def unwrap_tuple(value):
"""Unwrap single-item tuples to their contained value."""
return value[0] if isinstance(value, tuple) and len(value) == 1 else value