diff --git a/nodes/lora.py b/nodes/lora.py index 6a05db7..251e607 100644 --- a/nodes/lora.py +++ b/nodes/lora.py @@ -174,12 +174,12 @@ class Sage_CheckLorasForUpdates(ComfyNodeABC): pull_metadata(lora_path, timestamp=False, force=force) print(f"Update check complete for {lora[0]}") - if "update_available" in cache.data[lora_path]: - if cache.data[lora_path]["update_available"] == True: + if "update_available" in cache.by_path(lora_path): + if cache.by_path(lora_path)["update_available"] == True: print(f"Update found for {lora[0]}") lora_list.append(lora_path) - model_id = cache.data[lora_path]["modelId"] + model_id = cache.by_path(lora_path)["modelId"] latest_version = get_latest_model_version(model_id) latest_url = f"https://civitai.com/models/{model_id}?modelVersionId={latest_version}" lora_url_list.append(latest_url) diff --git a/nodes/metadata.py b/nodes/metadata.py index 972fcfc..5ea84f0 100644 --- a/nodes/metadata.py +++ b/nodes/metadata.py @@ -59,7 +59,7 @@ class Sage_ConstructMetadata(ComfyNodeABC): if lora_data != {}: resource_hashes.append(lora_data) - lora_hash = cache.data[lora_path]["hash"] + lora_hash = cache.hash[lora_path] lora_hashes += [f"{lora_name}: {lora_hash}"] lora_hash_string = "Lora hashes: " + ",".join(lora_hashes) diff --git a/nodes/model.py b/nodes/model.py index 78b02e4..1b4d174 100644 --- a/nodes/model.py +++ b/nodes/model.py @@ -40,7 +40,7 @@ class Sage_CheckpointLoaderRecent(ComfyNodeABC): model_info = { "type": "CKPT", "path": folder_paths.get_full_path_or_raise("checkpoints", ckpt_name) } pull_metadata(model_info["path"], True) - model_info["hash"] = cache.data[model_info["path"]]["hash"] + model_info["hash"] = cache.hash[model_info["path"]] model, clip, vae = loaders.checkpoint(model_info["path"]) result = (model, clip, vae, model_info) @@ -72,7 +72,7 @@ class Sage_CheckpointLoaderSimple(CheckpointLoaderSimple): model_info = { "type": "CKPT", "path": folder_paths.get_full_path_or_raise("checkpoints", ckpt_name) } pull_metadata(model_info["path"], True) - model_info["hash"] = cache.data[model_info["path"]]["hash"] + model_info["hash"] = cache.hash[model_info["path"]] model, clip, vae = loaders.checkpoint(model_info["path"]) return (model, clip, vae, model_info) @@ -99,7 +99,7 @@ class Sage_UNETLoader(UNETLoader): "path": folder_paths.get_full_path_or_raise("diffusion_models", unet_name) } pull_metadata(model_info["path"], True) - model_info["hash"] = cache.data[model_info["path"]]["hash"] + model_info["hash"] = cache.hash[model_info["path"]] return (loaders.unet(model_info["path"], weight_dtype), model_info) class Sage_CheckpointSelector(ComfyNodeABC): @@ -123,7 +123,7 @@ class Sage_CheckpointSelector(ComfyNodeABC): def get_checkpoint_info(self, ckpt_name) -> tuple: model_info = { "type": "CKPT", "path": folder_paths.get_full_path_or_raise("checkpoints", ckpt_name) } pull_metadata(model_info["path"], True) - model_info["hash"] = cache.data[model_info["path"]]["hash"] + model_info["hash"] = cache.hash[model_info["path"]] return (model_info,) class Sage_MultiModelPicker(ComfyNodeABC): @@ -175,7 +175,11 @@ class Sage_CacheMaintenance(ComfyNodeABC): DESCRIPTION = "Lets you remove entries for models that are no longer there. dup_hash returns a list of files with the same hash, and dup_model returns ones with the same civitai model id (but not neccessarily the same version)." def cache_maintenance(self, remove_ghost_entries) -> tuple[str, str, str, str]: - ghost_entries = [path for path in cache.data if not pathlib.Path(path).is_file()] + ghost_entries = [] + for key in cache.hash: + if not pathlib.Path(key).is_file(): + ghost_entries.append(key) + cache_by_hash = {} cache_by_id = {} dup_hash = {} @@ -183,15 +187,21 @@ class Sage_CacheMaintenance(ComfyNodeABC): not_on_civitai = [] out_of_date = [] - for model_path, data in cache.data.items(): - if 'hash' in data: - cache_by_hash.setdefault(data['hash'], []).append(model_path) - if 'modelId' in data: - cache_by_id.setdefault(data['modelId'], []).append(model_path) + for model_path, model_hash in cache.hash.items(): + if model_hash not in cache_by_hash: + cache_by_hash[model_hash] = [] + cache_by_hash[model_hash].append(model_path) + model_info = cache.by_path(model_path) + model_id = model_info.get("modelId", None) + if model_id: + if model_id not in cache_by_id: + cache_by_id[model_id] = [] + cache_by_id[model_id].append(model_path) + if remove_ghost_entries: for ghost in ghost_entries: - cache.data.pop(ghost) + cache.hash.pop(ghost) cache.save() dup_hash = {h: paths for h, paths in cache_by_hash.items() if len(paths) > 1} @@ -200,11 +210,13 @@ class Sage_CacheMaintenance(ComfyNodeABC): dup_hash_json = json.dumps(dup_hash, separators=(",", ":"), sort_keys=True, indent=4) dup_id_json = json.dumps(dup_id, separators=(",", ":"), sort_keys=True, indent=4) - for model_path, data in cache.data.items(): - if data.get("civitai", "False") == "False": + for model_path, model_hash in cache.hash.items(): + model_info = cache.by_path(model_path) + in_civitai = model_info['civitai'] + if in_civitai != True: not_on_civitai.append(model_path) - if data.get("update_available", "False") == "True": + if model_info.get("update_available", False): out_of_date.append(model_path) not_on_civitai_str = str(not_on_civitai) @@ -213,7 +225,7 @@ class Sage_CacheMaintenance(ComfyNodeABC): class Sage_ModelReport(ComfyNodeABC): @classmethod - def INPUT_TYPES(cls) -> InputTypeDict: + def INPUT_TYPES(cls): return { "required": { "scan_models": (["none", "loras", "checkpoints", "all"], {"defaultInput": False, "default": "none"}), @@ -239,8 +251,6 @@ class Sage_ModelReport(ComfyNodeABC): the_checkpoint_paths = folder_paths.get_folder_paths("checkpoints") the_paths = [*the_lora_paths, *the_checkpoint_paths] - print(f"Scanning {len(the_paths)} paths.") - print(f"the_paths == {the_paths}") if the_paths != []: model_scan(the_paths, force=force_recheck) def pull_list(self, scan_models, force_recheck) -> tuple[str, str]: @@ -251,8 +261,8 @@ class Sage_ModelReport(ComfyNodeABC): self.get_files(scan_models, force_recheck) - for model_path in cache.data.keys(): - cur = cache.data.get(model_path, {}) + for model_path in cache.hash.keys(): + cur = cache.info.get(cache.hash[model_path], {}) baseModel = cur.get('baseModel', None) if cur.get('model', {}).get('type', None) == "Checkpoint": if baseModel not in sorted_models: sorted_models[baseModel] = [] diff --git a/nodes/text.py b/nodes/text.py index 28e4c08..1e8f544 100644 --- a/nodes/text.py +++ b/nodes/text.py @@ -127,7 +127,7 @@ class Sage_ViewText(ComfyNodeABC): DEPRECATED = True def show_text(self, text) -> tuple[str]: - print(f"String is '{text}'") + #print(f"String is '{text}'") return { "ui": {"text": text}, "result" : (text,) } @@ -150,15 +150,15 @@ class Sage_ViewAnything(ComfyNodeABC): INPUT_IS_LIST = True def show_text(self, any) -> dict: - print(f"Text is '{any}'") + #print(f"Text is '{any}'") str = "" if isinstance(any, list): for t in any: str += f"{t}\n" - print(f"String is '{t}'") + #print(f"String is '{t}'") else: str = any - print(f"String is '{str}'") + #print(f"String is '{str}'") return { "ui": {"text": str}, "result" : (str,) } class Sage_PonyPrefix(ComfyNodeABC): diff --git a/nodes/util.py b/nodes/util.py index f959873..2976d35 100644 --- a/nodes/util.py +++ b/nodes/util.py @@ -203,10 +203,10 @@ class Sage_GetFileHash(ComfyNodeABC): try: file_path = folder_paths.get_full_path_or_raise(base_dir, filename) pull_metadata(file_path) - the_hash = cache.data[file_path]["hash"] + the_hash = cache.hash[file_path] except: - print(f"Unable to hash file '{file_path}'. \n") + print(f"Unable to hash file '{filename}'. \n") the_hash = "" - print(f"Hash for '{file_path}': {the_hash}") + print(f"Hash for '{filename}': {the_hash}") return (str(the_hash),) diff --git a/utils/cache.py b/utils/cache.py index 4c4694f..03e2c06 100644 --- a/utils/cache.py +++ b/utils/cache.py @@ -4,45 +4,85 @@ import pathlib import folder_paths -base_path = pathlib.Path(os.path.dirname(os.path.realpath(__file__))).parent -#print(f"Loading SageUtils cache from {str(base_path)}") - users_path = pathlib.Path(folder_paths.get_user_directory()) sage_users_path = users_path / "default" / "SageUtils" os.makedirs(str(sage_users_path), exist_ok=True) class SageCache: - def __init__(self, path): + def __init__(self): if not (sage_users_path / "sage_cache.json").is_file(): print("No cache file found in user directory.") - if (pathlib.Path(path) / "sage_cache.json").is_file(): - with open((pathlib.Path(path) / "sage_cache.json"), "r") as read_file: - temp = json.load(read_file) - with open((sage_users_path / "sage_cache.json"), "w") as write_file: - json.dump(temp, write_file, separators=(",", ":"), sort_keys=True, indent=4) - - print("Copied old cache file to {str(sage_users_path)}.") - - self.path = sage_users_path / "sage_cache.json" + self.main_path = sage_users_path / "sage_cache.json" + self.info_path = sage_users_path / "sage_cache_info.json" + self.hash_path = sage_users_path / "sage_cache_hash.json" self.data = {} + self.hash = {} + self.info = {} + def by_path(self, file_path): + the_hash = self.hash.get(file_path, "") + if the_hash: + return self.info.get(the_hash, {}) + else: + print(f"No hash found for file: {file_path}") + return {} + + def by_hash(self, file_hash): + return self.info.get(file_hash, {}) + + def convert_old_cache(self): + print("Converting old cache format to new format.") + for key in self.data: + current_hash = self.data[key].get("hash", "") + if current_hash: + self.hash[key] = current_hash + + # Add the ones not on civitai first + if self.data[key].get("civitai", False) == False: + self.info[current_hash] = self.data[key] + + for key in self.data: + current_hash = self.data[key].get("hash", "") + # Add the ones on civitai, overwriting the previous ones + if current_hash and self.data[key].get("civitai", False): + self.info[current_hash] = self.data[key] + def load(self): try: - if self.path.is_file(): - with self.path.open("r") as read_file: - self.data = json.load(read_file) + if self.hash_path.is_file() and self.info_path.is_file(): + with self.hash_path.open("r") as read_file: + self.hash = json.load(read_file) + with self.info_path.open("r") as read_file: + self.info = json.load(read_file) + else: + if self.main_path.is_file(): + with self.main_path.open("r") as read_file: + self.data = json.load(read_file) + self.convert_old_cache() + except Exception as e: print(f"Unable to load cache: {e}") + def save(self): try: if self.data: - with self.path.open("w") as output_file: + with self.main_path.open("w") as output_file: json.dump(self.data, output_file, separators=(",", ":"), sort_keys=True, indent=4) else: print("Skipping saving cache, as the cache is empty.") + if self.hash: + with self.hash_path.open("w") as output_file: + json.dump(self.hash, output_file, separators=(",", ":"), sort_keys=True, indent=4) + else: + print("Skipping saving hash, as the hash is empty.") + if self.info: + with self.info_path.open("w") as output_file: + json.dump(self.info, output_file, separators=(",", ":"), sort_keys=True, indent=4) + else: + print("Skipping saving info, as the info is empty.") except Exception as e: print(f"Unable to save cache: {e}") -cache = SageCache(base_path) +cache = SageCache() diff --git a/utils/helpers.py b/utils/helpers.py index 98d6247..9bc5c38 100644 --- a/utils/helpers.py +++ b/utils/helpers.py @@ -70,12 +70,12 @@ def get_civitai_model_json(modelId): def get_model_info(lora_path, weight = None): ret = {} try: - ret["type"] = cache.data[lora_path]["model"]["type"] + ret["type"] = cache.by_path(lora_path)["model"]["type"] if (ret["type"] == "LORA") and (weight is not None): ret["weight"] = weight - ret["modelVersionId"] = cache.data[lora_path]["id"] - ret["modelName"] = cache.data[lora_path]["model"]["name"] - ret["modelVersionName"] = cache.data[lora_path]["name"] + ret["modelVersionId"] = cache.by_path(lora_path)["id"] + ret["modelName"] = cache.by_path(lora_path)["model"]["name"] + ret["modelVersionName"] = cache.by_path(lora_path)["name"] except: ret = {} return ret @@ -108,9 +108,9 @@ def get_file_sha256(path): def last_used(file_path): cache.load() - - if file_path in cache.data: - last_used = cache.data[file_path].get("lastUsed", None) + + 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: @@ -144,11 +144,11 @@ def pull_metadata(file_path, timestamp = True, force = False): cache.load() print(f"Pull metadata for {file_path}.") - hash = cache.data.get(file_path, {}).get("hash", "") + hash = cache.hash.get(file_path, "") if not hash: - cache.data[file_path] = {"hash": get_file_sha256(file_path)} - hash = cache.data[file_path]["hash"] + hash = get_file_sha256(file_path) + cache.hash[file_path] = hash else: time.sleep(2) @@ -156,7 +156,7 @@ def pull_metadata(file_path, timestamp = True, force = False): check_recent = False metadata_days_recheck = 0 hash_recheck = 30 - file_cache = cache.data.get(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 if last_used_date is not None: @@ -182,7 +182,7 @@ def pull_metadata(file_path, timestamp = True, force = False): if new_hash != hash: print(f"Hash mismatch. Pulling new hash.") - cache.data[file_path]['hash'] = new_hash + cache.hash[file_path] = new_hash json = get_civitai_model_version_json_by_hash(new_hash) if 'error' in json and 'modelId' in file_cache: @@ -229,6 +229,7 @@ def pull_metadata(file_path, timestamp = True, force = False): file_cache['lastUsed'] = datetime.datetime.now().isoformat() cache.data[file_path] = file_cache + cache.info[hash] = file_cache cache.save() def lora_to_string(lora_name, model_weight, clip_weight): @@ -249,7 +250,7 @@ def get_lora_hash(lora_name): lora_path = folder_paths.get_full_path_or_raise("loras", lora_name) pull_metadata(lora_path) - return cache.data[lora_path]["hash"] + return cache.hash[lora_path] def model_scan(the_path, force = False): the_paths = the_path @@ -297,13 +298,13 @@ def get_recently_used_models(model_type): full_model_list = folder_paths.get_filename_list(model_type) for item in full_model_list: model_path = folder_paths.get_full_path_or_raise(model_type, item) - if model_path not in cache.data.keys(): + if model_path not in cache.hash.keys(): continue - if 'lastUsed' not in cache.data[model_path]: + if 'lastUsed' not in cache.by_path(model_path): continue - last = cache.data[model_path]['lastUsed'] + last = cache.by_path(model_path)['lastUsed'] last_used = datetime.datetime.fromisoformat(last) #print(f"{model_path} - last: {last} last_used: {last_used}") if (datetime.datetime.now() - last_used).days <= 7: diff --git a/utils/lora_stack.py b/utils/lora_stack.py index d662215..a165377 100644 --- a/utils/lora_stack.py +++ b/utils/lora_stack.py @@ -4,10 +4,10 @@ from .helpers import pull_metadata, clean_keywords def get_lora_keywords(lora_name): lora_path = folder_paths.get_full_path_or_raise("loras", lora_name) - if cache.data.get(lora_path, {}).get("trainedWords", None) is None: + if cache.by_path(lora_path).get("trainedWords", None) is None: pull_metadata(lora_path, True) - return cache.data.get(lora_path, {}).get("trainedWords", []) + return cache.by_path(lora_path).get("trainedWords", []) def get_lora_stack_keywords(lora_stack = None): lora_keywords = []