Split Loader.py

This commit is contained in:
Nourepide
2023-12-07 20:45:23 +03:00
parent 89aba77199
commit 759ac287fb
12 changed files with 495 additions and 381 deletions
+1
View File
@@ -320,5 +320,6 @@ pip-selfcheck.json
### Allor ###
config.json
resources/timestamp.json
resources/logs/*
# End of https://www.toptal.com/developers/gitignore/api/python,venv,visualstudiocode,pycharm
-371
View File
@@ -1,371 +0,0 @@
import json
import os
import platform
import time
from pathlib import Path
import folder_paths
import nodes
class Loader:
def __init__(self):
pass
__ROOT_PATH = os.path.dirname(os.path.abspath(__file__))
__TEMPLATE_PATH = os.path.join(__ROOT_PATH, "resources/template.json")
__TIMESTAMP_PATH = os.path.join(__ROOT_PATH, "resources/timestamp.json")
__CONFIG_PATH = os.path.join(__ROOT_PATH, "config.json")
__GIT_PATH = Path(os.path.join(__ROOT_PATH, ".git"))
__DAY_SECONDS = 24 * 60 * 60
__WEEK_SECONDS = 7 * __DAY_SECONDS
__MONTH_SECONDS = 30 * __DAY_SECONDS
def __log(self, text):
print("\033[92m[Allor]\033[0m: " + text)
def __error(self, text):
print("\033[91m[Allor]\033[0m: " + text)
def __notification(self, text):
print("\033[94m[Allor]\033[0m: " + text)
def __new_line(self):
print()
def __warning_unstable_branch(self):
self.__new_line()
self.__error("Attention! You are currently using an unstable \"main\" update branch intended for the development of Allor 2.")
self.__error("Please be aware that changes made in Allor 2 may disrupt your current workflow.")
self.__error("Nodes may be renamed, parameters within them may be altered or even removed.")
self.__new_line()
self.__error("If backward compatibility of your workflow is important to you, "
"you can change the \"branch_name\" parameter to \"allor-1\" in your config.json.")
self.__error("Switch the \"confirm_unstable_agreement\" parameter in your config.json to \"true\", "
"if you are prepared for potential changes and are willing to modify your current workflow from time to time.")
self.__error("This will result in this warning no longer appearing.")
self.__new_line()
self.__notification("We appreciate your support and understanding during this transition period.")
self.__notification("Thank you for using Allor 2.\n")
def __create_config(self):
with open(self.__CONFIG_PATH, "w", encoding="utf-8") as f:
json.dump(self.__template(), f, ensure_ascii=False, indent=4)
def __create_timestamp(self):
with open(self.__TIMESTAMP_PATH, "w", encoding="utf-8") as f:
json.dump({"timestamp": 0}, f, ensure_ascii=False, indent=4)
def __get_template(self):
with open(self.__TEMPLATE_PATH, "r") as f:
template = json.load(f)
if "__comment" in template:
del template["__comment"]
return template
def __get_config(self):
with open(self.__CONFIG_PATH, "r") as f:
return json.load(f)
def __get_timestamp(self):
with open(self.__TIMESTAMP_PATH, "r") as f:
return json.load(f)
def __update_config(self, template, source):
def update_source(__template, __source):
for k, v in __template.items():
if k not in __source:
if isinstance(v, dict):
__source[k] = {}
else:
__source[k] = v
if isinstance(v, dict):
__source[k] = update_source(v, __source[k])
return __source
def delete_keys(__template, __source):
keys_to_delete = [k for k in __source if k not in __template]
for k in keys_to_delete:
del __source[k]
return __source
def sync_order(__template, __source):
new_source = {}
for key in __template:
if key in __source:
if isinstance(__template[key], dict):
new_source[key] = sync_order(__template[key], __source[key])
else:
new_source[key] = __source[key]
return new_source
source = update_source(template, source)
source = delete_keys(template, source)
source = sync_order(template, source)
with open(self.__CONFIG_PATH, "w", encoding="utf-8") as f:
json.dump(source, f, ensure_ascii=False, indent=4)
def __update_timestamp(self):
with open(self.__TIMESTAMP_PATH, "w", encoding="utf-8") as f:
json.dump({"timestamp": time.time()}, f, ensure_ascii=False, indent=4)
__template = __get_template
__config = __get_config
__timestamp = __get_timestamp
def __get_fonts_folder_path(self):
system = platform.system()
user_home = os.path.expanduser('~')
config_font_path = os.path.join(folder_paths.base_path, *self.__config()["fonts"]["folder_path"].replace("\\", "/").split("/"))
if not os.path.exists(config_font_path):
os.makedirs(config_font_path, exist_ok=True)
paths = [config_font_path]
if self.__config()["fonts"]["system_fonts"]:
if system == "Windows":
paths.append(os.path.join(os.environ["WINDIR"], "Fonts"))
elif system == "Darwin":
paths.append(os.path.join("/Library", "Fonts"))
elif system == "Linux":
paths.append(os.path.join("/usr", "share", "fonts"))
paths.append(os.path.join("/usr", "local", "share", "fonts"))
if self.__config()["fonts"]["user_fonts"]:
if system == "Darwin":
paths.append(os.path.join(user_home, "Library", "Fonts"))
elif system == "Linux":
paths.append(os.path.join(user_home, ".fonts"))
return [path for path in paths if os.path.exists(path)]
def __get_keys(self, json_obj, prefix=''):
keys = []
for k, v in json_obj.items():
if isinstance(v, dict):
keys.extend(self.__get_keys(v, prefix + k + '.'))
else:
keys.append(prefix + k)
return set(keys)
def __check_json_keys(self, json1, json2):
keys1 = self.__get_keys(json1)
keys2 = self.__get_keys(json2)
return keys1 == keys2
def setup_config(self):
if not os.path.exists(self.__CONFIG_PATH):
self.__log("Creating config.json")
self.__create_config()
else:
if not self.__check_json_keys(self.__template(), self.__config()):
self.__log("Updating config.json")
self.__update_config(self.__template(), self.__config())
def setup_timestamp(self):
if not os.path.exists(self.__TIMESTAMP_PATH):
self.__log("Creating timestamp.json")
self.__create_timestamp()
def check_updates(self):
# confirm_unstable_agreement = self.__config()["updates"]["confirm_unstable_agreement"]
confirm_unstable_agreement = True
branch_name = self.__config()["updates"]["branch_name"]
update_frequency = self.__config()["updates"]["update_frequency"].lower()
valid_frequencies = ["always", "day", "week", "month", "never"]
time_difference = time.time() - self.__timestamp()["timestamp"]
if update_frequency == valid_frequencies[0]:
it_is_time_for_update = True
elif update_frequency == valid_frequencies[1]:
it_is_time_for_update = time_difference >= self.__DAY_SECONDS
elif update_frequency == valid_frequencies[2]:
it_is_time_for_update = time_difference >= self.__WEEK_SECONDS
elif update_frequency == valid_frequencies[3]:
it_is_time_for_update = time_difference >= self.__MONTH_SECONDS
elif update_frequency == valid_frequencies[4]:
it_is_time_for_update = False
else:
self.__error(f"Unknown update frequency - {update_frequency}, available: {valid_frequencies}")
return
if not confirm_unstable_agreement and branch_name == "main" and update_frequency != "never":
self.__warning_unstable_branch()
if it_is_time_for_update:
if not (self.__GIT_PATH.exists() or self.__GIT_PATH.is_dir()):
self.__error("Root directory of Allor is not a git repository. Update canceled.")
return
try:
import git
from git import Repo
from git import GitCommandError
# noinspection PyTypeChecker, PyUnboundLocalVariable
repo = Repo(self.__ROOT_PATH, odbt=git.db.GitDB)
current_commit = repo.head.commit.hexsha
repo.remotes.origin.fetch()
latest_commit = getattr(repo.remotes.origin.refs, branch_name).commit.hexsha
if current_commit == latest_commit:
if self.__config()["updates"]["notify_if_no_new_updates"]:
self.__notification("No new updates.")
else:
if self.__config()["updates"]["notify_if_has_new_updates"]:
self.__notification("New updates are available.")
if self.__config()["updates"]["auto_update"]:
update_mode = self.__config()["updates"]["update_mode"].lower()
valid_modes = ["soft", "hard"]
if repo.active_branch.name != branch_name:
try:
repo.git.checkout(branch_name)
except GitCommandError:
self.__error(f"An error occurred while switching to the branch {branch_name}.")
return
if update_mode == "soft":
try:
repo.git.pull()
except GitCommandError:
self.__error("An error occurred during the update. "
"It is recommended to use \"hard\" update mode. "
"But be careful, it erases all personal changes from Allor repository.")
elif update_mode == "hard":
repo.git.reset('--hard', 'origin/' + branch_name)
else:
self.__error(f"Unknown update mode - {update_mode}, available: {valid_modes}")
return
self.__notification("Update complete.")
self.__update_timestamp()
except ImportError:
self.__error("GitPython is not installed.")
def setup_rembg(self):
os.environ["U2NET_HOME"] = folder_paths.models_dir + "/onnx"
def setup_paths(self):
fonts_folder_path = self.__get_fonts_folder_path()
folder_paths.folder_names_and_paths["onnx"] = ([os.path.join(folder_paths.models_dir, "onnx")], {".onnx"})
folder_paths.folder_names_and_paths["fonts"] = (fonts_folder_path, {".otf", ".ttf"})
def setup_override(self):
override_nodes_len = 0
def override(function):
start_len = nodes.NODE_CLASS_MAPPINGS.__len__()
nodes.NODE_CLASS_MAPPINGS = dict(
filter(function, nodes.NODE_CLASS_MAPPINGS.items())
)
return start_len - nodes.NODE_CLASS_MAPPINGS.__len__()
if self.__config()["override"]["postprocessing"]:
override_nodes_len += override(lambda item: not item[1].CATEGORY.startswith("image/postprocessing"))
if self.__config()["override"]["transform"]:
override_nodes_len += override(lambda item: not item[0] == "ImageScale" and not item[0] == "ImageScaleBy" and not item[0] == "ImageInvert")
if self.__config()["override"]["debug"]:
nodes.VAEDecodeTiled.CATEGORY = "latent"
nodes.VAEEncodeTiled.CATEGORY = "latent"
override_nodes_len += override(lambda item: not item[1].CATEGORY.startswith("_for_testing"))
self.__log(str(override_nodes_len) + " standard nodes was overridden.")
def get_modules(self):
modules = dict()
if self.__config()["modules"]["AlphaChanel"]:
from .modules import AlphaChanel
modules.update(AlphaChanel.NODE_CLASS_MAPPINGS)
if self.__config()["modules"]["Clamp"]:
from .modules import Clamp
modules.update(Clamp.NODE_CLASS_MAPPINGS)
if self.__config()["modules"]["ImageBatch"]:
from .modules import ImageBatch
modules.update(ImageBatch.NODE_CLASS_MAPPINGS)
if self.__config()["modules"]["ImageComposite"]:
from .modules import ImageComposite
modules.update(ImageComposite.NODE_CLASS_MAPPINGS)
if self.__config()["modules"]["ImageContainer"]:
from .modules import ImageContainer
modules.update(ImageContainer.NODE_CLASS_MAPPINGS)
if self.__config()["modules"]["ImageDraw"]:
from .modules import ImageDraw
modules.update(ImageDraw.NODE_CLASS_MAPPINGS)
if self.__config()["modules"]["ImageEffects"]:
from .modules import ImageEffects
modules.update(ImageEffects.NODE_CLASS_MAPPINGS)
if self.__config()["modules"]["ImageFilter"]:
from .modules import ImageFilter
modules.update(ImageFilter.NODE_CLASS_MAPPINGS)
if self.__config()["modules"]["ImageNoise"]:
from .modules import ImageNoise
modules.update(ImageNoise.NODE_CLASS_MAPPINGS)
if self.__config()["modules"]["ImageSegmentation"]:
from .modules import ImageSegmentation
modules.update(ImageSegmentation.NODE_CLASS_MAPPINGS)
if self.__config()["modules"]["ImageText"]:
from .modules import ImageText
modules.update(ImageText.NODE_CLASS_MAPPINGS)
if self.__config()["modules"]["ImageTransform"]:
from .modules import ImageTransform
modules.update(ImageTransform.NODE_CLASS_MAPPINGS)
modules_len = dict(
filter(
lambda item: item[1],
self.__config()["modules"].items()
)
).__len__()
nodes_len = modules.__len__()
self.__log(str(modules_len) + " modules enabled.")
self.__log(str(nodes_len) + " nodes was loaded.")
return modules
+10 -9
View File
@@ -1,12 +1,13 @@
from .Loader import Loader
from .boot.Config import Config
from .boot.Update import Update
from .boot.Paths import Paths
from .boot.Override import Override
from .boot.Modules import Modules
loader = Loader()
config = Config().initiate()
loader.setup_config()
loader.setup_timestamp()
loader.check_updates()
loader.setup_rembg()
loader.setup_paths()
loader.setup_override()
Update(config).initiate()
Paths(config).initiate()
Override(config).initiate()
NODE_CLASS_MAPPINGS = loader.get_modules()
NODE_CLASS_MAPPINGS = Modules(config).initiate()
+98
View File
@@ -0,0 +1,98 @@
import json
from .Constants import Constants
from .Logger import Logger
class Config:
def __init__(self):
self.__logger = Logger()
self.__template = self.__get_template()
if not Constants.CONFIG_PATH.exists():
self.__logger.info("Creating configuration file.")
self.__create_config()
self.__config = self.__get_config()
def initiate(self):
if not self.__verify_keys(self.__template, self.__config):
self.__logger.info("Updating configuration file.")
self.__update_config(self.__template, self.__config)
return self.__get_config()
def __create_config(self):
with open(Constants.CONFIG_PATH, "w", encoding="utf-8") as f:
json.dump(self.__template, f, ensure_ascii=False, indent=4)
def __get_template(self):
with open(Constants.TEMPLATE_PATH, "r") as f:
template = json.load(f)
if "__comment" in template:
del template["__comment"]
return template
def __get_config(self):
with open(Constants.CONFIG_PATH, "r") as f:
return json.load(f)
def __verify_keys(self, json1, json2):
def get_keys(json_obj, prefix=''):
keys = []
for k, v in json_obj.items():
if isinstance(v, dict):
keys.extend(get_keys(v, prefix + k + '.'))
else:
keys.append(prefix + k)
return set(keys)
keys1 = get_keys(json1)
keys2 = get_keys(json2)
return keys1 == keys2
def __update_config(self, template, source):
def update_source(__template, __source):
for k, v in __template.items():
if k not in __source:
if isinstance(v, dict):
__source[k] = {}
else:
__source[k] = v
if isinstance(v, dict):
__source[k] = update_source(v, __source[k])
return __source
def delete_keys(__template, __source):
keys_to_delete = [k for k in __source if k not in __template]
for k in keys_to_delete:
del __source[k]
return __source
def sync_order(__template, __source):
new_source = {}
for key in __template:
if key in __source:
if isinstance(__template[key], dict):
new_source[key] = sync_order(__template[key], __source[key])
else:
new_source[key] = __source[key]
return new_source
source = update_source(template, source)
source = delete_keys(template, source)
source = sync_order(template, source)
with open(Constants.CONFIG_PATH, "w", encoding="utf-8") as f:
json.dump(source, f, ensure_ascii=False, indent=4)
+16
View File
@@ -0,0 +1,16 @@
import os
from datetime import datetime
from pathlib import Path
class Constants:
ROOT_PATH = Path(os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
TEMPLATE_PATH = ROOT_PATH / "resources/template.json"
TIMESTAMP_PATH = ROOT_PATH / "resources/timestamp.json"
CONFIG_PATH = ROOT_PATH / "config.json"
GIT_PATH = ROOT_PATH / ".git"
LOG_PATH = ROOT_PATH / f"resources/logs/log_{datetime.now().strftime('%Y_%m_%d_%H_%M_%S')}.txt"
DAY_SECONDS = 24 * 60 * 60
WEEK_SECONDS = 7 * DAY_SECONDS
MONTH_SECONDS = 30 * DAY_SECONDS
+98
View File
@@ -0,0 +1,98 @@
import glob
import os
import re
from .Constants import Constants
class Logger:
file = Constants.LOG_PATH
levels = {
"FATAL": 92,
"ERROR": 91,
"WARN": 93,
"INFO": 94,
"DEBUG": 95,
"TRACE": 96,
"LINE": 0
}
def __init__(self):
pass
def __log(self, text, level):
if level == "LINE":
print()
else:
print(f"\033[{self.levels.get(level, '')}m[Allor]\033[0m: " + text)
if Logger.file:
directory = Logger.file.parent
if not directory.exists():
directory.mkdir(parents=True)
if not Logger.file.exists():
log_files = list(glob.glob(str(directory / 'log_*.txt')))
if len(log_files) > 5:
log_files.sort(key=os.path.getmtime)
os.remove(log_files[0])
Logger.file.touch()
with open(Logger.file, "a") as f:
if level == "LINE":
f.write("\n")
else:
text = re.compile(r'\x1B(?:[@-Z\\-_]|\[[0-?]*[ -/]*[@-~])').sub("", text)
f.write(f"[{level}]: {text}\n")
def fatal(self, text):
self.__log(text, "FATAL")
def error(self, text):
self.__log(text, "ERROR")
def warn(self, text):
self.__log(text, "WARN")
def info(self, text):
self.__log(text, "INFO")
def debug(self, text):
self.__log(text, "DEBUG")
def trace(self, text):
self.__log(text, "TRACE")
def line(self, count=1):
self.__log("\n" * count, "LINE")
def warning_unstable_branch(self, branch_name="main"):
warn_messages = [
f"Attention! You are currently using an unstable \033[1m{branch_name}\033[0m update branch.",
f"This branch is intended for the development of \033[1mAllor v.2\033[0m.",
f"Please be aware that changes made in \033[1mAllor v.2\033[0m may disrupt your current workflow.",
"Nodes may be renamed, parameters within them may be altered or even removed.",
"If backward compatibility of your workflow is important to you, consider this.",
f"You can change the \033[1m\"branch_name\"\033[0m parameter to \033[1m\"v.1\"\033[0m in your \033[1mconfig.json.\033[0m",
"If you are prepared for potential changes, you can modify your current workflow.",
f"To accept, switch the \033[1m\"confirm_unstable\"\033[0m parameter in your \033[1mconfig.json.\033[0m to \033[1m\"true\".\033[0m",
"This will result in this warning no longer appearing."
]
info_messages = [
"We appreciate your support and understanding during this transition period.",
f"Thank you and welcome to \033[1mAllor v.2\033[0m.\n"
]
for message in warn_messages:
self.warn(message)
self.line()
for message in info_messages:
self.info(message)
+72
View File
@@ -0,0 +1,72 @@
from .Logger import Logger
class Modules:
def __init__(self, config):
self.__logger = Logger()
self.__config = config
def initiate(self):
modules = dict()
if self.__config["modules"]["AlphaChanel"]:
from ..modules import AlphaChanel
modules.update(AlphaChanel.NODE_CLASS_MAPPINGS)
if self.__config["modules"]["Clamp"]:
from ..modules import Clamp
modules.update(Clamp.NODE_CLASS_MAPPINGS)
if self.__config["modules"]["ImageBatch"]:
from ..modules import ImageBatch
modules.update(ImageBatch.NODE_CLASS_MAPPINGS)
if self.__config["modules"]["ImageComposite"]:
from ..modules import ImageComposite
modules.update(ImageComposite.NODE_CLASS_MAPPINGS)
if self.__config["modules"]["ImageContainer"]:
from ..modules import ImageContainer
modules.update(ImageContainer.NODE_CLASS_MAPPINGS)
if self.__config["modules"]["ImageDraw"]:
from ..modules import ImageDraw
modules.update(ImageDraw.NODE_CLASS_MAPPINGS)
if self.__config["modules"]["ImageEffects"]:
from ..modules import ImageEffects
modules.update(ImageEffects.NODE_CLASS_MAPPINGS)
if self.__config["modules"]["ImageFilter"]:
from ..modules import ImageFilter
modules.update(ImageFilter.NODE_CLASS_MAPPINGS)
if self.__config["modules"]["ImageNoise"]:
from ..modules import ImageNoise
modules.update(ImageNoise.NODE_CLASS_MAPPINGS)
if self.__config["modules"]["ImageSegmentation"]:
from ..modules import ImageSegmentation
modules.update(ImageSegmentation.NODE_CLASS_MAPPINGS)
if self.__config["modules"]["ImageText"]:
from ..modules import ImageText
modules.update(ImageText.NODE_CLASS_MAPPINGS)
if self.__config["modules"]["ImageTransform"]:
from ..modules import ImageTransform
modules.update(ImageTransform.NODE_CLASS_MAPPINGS)
modules_len = dict(
filter(
lambda item: item[1],
self.__config["modules"].items()
)
).__len__()
nodes_len = modules.__len__()
self.__logger.info(str(modules_len) + " modules were enabled.")
self.__logger.info(str(nodes_len) + " nodes were loaded.")
return modules
+34
View File
@@ -0,0 +1,34 @@
import nodes
from .Logger import Logger
class Override:
def __init__(self, config):
self.__logger = Logger()
self.__config = config
def initiate(self):
override_nodes_len = 0
def override(function):
start_len = nodes.NODE_CLASS_MAPPINGS.__len__()
nodes.NODE_CLASS_MAPPINGS = dict(
filter(function, nodes.NODE_CLASS_MAPPINGS.items())
)
return start_len - nodes.NODE_CLASS_MAPPINGS.__len__()
if self.__config["override"]["postprocessing"]:
override_nodes_len += override(lambda item: not item[1].CATEGORY.startswith("image/postprocessing"))
if self.__config["override"]["transform"]:
override_nodes_len += override(lambda item: not item[0] == "ImageScale" and not item[0] == "ImageScaleBy" and not item[0] == "ImageInvert")
if self.__config["override"]["debug"]:
nodes.VAEDecodeTiled.CATEGORY = "latent"
nodes.VAEEncodeTiled.CATEGORY = "latent"
override_nodes_len += override(lambda item: not item[1].CATEGORY.startswith("_for_testing"))
self.__logger.info(str(override_nodes_len) + " nodes were overridden.")
+48
View File
@@ -0,0 +1,48 @@
import os
import platform
import folder_paths
from .Logger import Logger
class Paths:
def __init__(self, config):
self.__logger = Logger()
self.__config = config
def initiate(self):
fonts_folder_path = self.__get_fonts_folder_path()
os.environ["U2NET_HOME"] = folder_paths.models_dir + "/onnx"
folder_paths.folder_names_and_paths["onnx"] = ([os.path.join(folder_paths.models_dir, "onnx")], {".onnx"})
folder_paths.folder_names_and_paths["fonts"] = (fonts_folder_path, {".otf", ".ttf"})
def __get_fonts_folder_path(self):
system = platform.system()
user_home = os.path.expanduser('~')
config_font_path = os.path.join(folder_paths.base_path, *self.__config["fonts"]["folder_path"].replace("\\", "/").split("/"))
if not os.path.exists(config_font_path):
os.makedirs(config_font_path, exist_ok=True)
paths = [config_font_path]
if self.__config["fonts"]["system_fonts"]:
if system == "Windows":
paths.append(os.path.join(os.environ["WINDIR"], "Fonts"))
elif system == "Darwin":
paths.append(os.path.join("/Library", "Fonts"))
elif system == "Linux":
paths.append(os.path.join("/usr", "share", "fonts"))
paths.append(os.path.join("/usr", "local", "share", "fonts"))
if self.__config["fonts"]["user_fonts"]:
if system == "Darwin":
paths.append(os.path.join(user_home, "Library", "Fonts"))
elif system == "Linux":
paths.append(os.path.join(user_home, ".fonts"))
return [path for path in paths if os.path.exists(path)]
+117
View File
@@ -0,0 +1,117 @@
import json
import time
from .Constants import Constants
from .Logger import Logger
class Update:
def __init__(self, config):
self.__logger = Logger()
self.__config = config
if not Constants.CONFIG_PATH.exists():
self.__logger.info("Creating timestamp file.")
self.__create_timestamp()
self.__timestamp = self.__get_timestamp()
def initiate(self):
confirm_unstable_agreement = self.__config["updates"]["confirm_unstable"]
branch_name = self.__config["updates"]["branch_name"]
update_frequency = self.__config["updates"]["update_frequency"].lower()
time_difference = time.time() - self.__timestamp["timestamp"]
valid_frequencies = {
"always": True,
"day": time_difference >= Constants.DAY_SECONDS,
"week": time_difference >= Constants.WEEK_SECONDS,
"month": time_difference >= Constants.MONTH_SECONDS,
"never": False
}
try:
it_is_time_for_update = valid_frequencies[update_frequency]
except KeyError:
self.__logger.error(f"Unknown update frequency - {update_frequency}, available: {list(valid_frequencies.keys())}")
return
if not confirm_unstable_agreement and branch_name == "main" and update_frequency != "never":
self.__logger.warning_unstable_branch(branch_name)
if it_is_time_for_update:
if not (Constants.GIT_PATH.exists() and Constants.GIT_PATH.is_dir()):
self.__logger.error("Root directory of Allor is not a git repository. Update canceled.")
return
self.__update_allor(branch_name)
def __get_timestamp(self):
with open(Constants.TIMESTAMP_PATH, "r") as f:
return json.load(f)
def __create_timestamp(self):
with open(Constants.TIMESTAMP_PATH, "w", encoding="utf-8") as f:
json.dump({"timestamp": 0}, f, ensure_ascii=False, indent=4)
def __update_timestamp(self):
with open(Constants.TIMESTAMP_PATH, "w", encoding="utf-8") as f:
json.dump({"timestamp": time.time()}, f, ensure_ascii=False, indent=4)
def __update_allor(self, branch_name):
try:
import git
from git import Repo
from git import GitCommandError
# noinspection PyTypeChecker, PyUnboundLocalVariable
repo = Repo(Constants.ROOT_PATH, odbt=git.db.GitDB)
current_commit = repo.head.commit.hexsha
repo.remotes.origin.fetch()
latest_commit = getattr(repo.remotes.origin.refs, branch_name).commit.hexsha
if current_commit == latest_commit:
if self.__config["updates"]["notify_if_no_new_updates"]:
self.__logger.info("No new updates.")
else:
if self.__config["updates"]["notify_if_has_new_updates"]:
self.__logger.info("New updates are available.")
if self.__config["updates"]["auto_update"]:
update_mode = self.__config["updates"]["update_mode"].lower()
valid_modes = ["soft", "hard"]
if repo.active_branch.name != branch_name:
try:
repo.git.checkout(branch_name)
except GitCommandError:
self.__logger.error(f"An error occurred while switching to the branch {branch_name}.")
return
if update_mode == "soft":
try:
repo.git.pull()
except GitCommandError:
self.__logger.error("An error occurred during the update. "
"It is recommended to use \"hard\" update mode. "
"But be careful, it erases all personal changes from Allor repository.")
elif update_mode == "hard":
repo.git.reset('--hard', 'origin/' + branch_name)
else:
self.__logger.error(f"Unknown update mode - {update_mode}, available: {valid_modes}")
return
self.__logger.info("Update complete.")
self.__update_timestamp()
except ImportError:
self.__logger.error("GitPython is not installed.")
View File
+1 -1
View File
@@ -26,7 +26,7 @@
"auto_update": true,
"branch_name": "main",
"update_mode": "soft",
"confirm_unstable_agreement": false
"confirm_unstable": false
},
"fonts": {
"folder_path": "comfy_extras/fonts",