diff --git a/ppp.py b/ppp.py index 7f3cc19..c2094cf 100644 --- a/ppp.py +++ b/ppp.py @@ -2,6 +2,7 @@ import csv import dataclasses from datetime import datetime from enum import Enum +from io import StringIO import json import logging from pathlib import Path @@ -11,7 +12,7 @@ import time from typing import Any, Callable, Optional import lark import numpy as np -import yaml +from ruamel.yaml import YAML as _YAML from pydantic import ValidationError from ppp_classes import ( @@ -275,9 +276,10 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in """Loads config files, performs model detection, and returns the resolved host config.""" main_folder = Path(__file__).resolve().parent default_config_file = str(main_folder / "ppp_config.yaml.defaults") + _yaml_rt = _YAML() try: with open(default_config_file, "r", encoding="utf-8") as f: - default_raw: dict[str, Any] = yaml.safe_load(f) + default_raw: dict[str, Any] = _yaml_rt.load(f) except Exception as exc: # pylint: disable=broad-exception-caught self.config = {} raise PPPInterrupt( @@ -295,7 +297,6 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in if isinstance(user_config_file, dict): user_cfg, _ = self.__parse_configuration(user_config_file, "forced configuration") else: - user_raw: dict[str, Any] = {} if user_config_file == "": if app == SUPPORTED_APPS.comfyui.value: try: @@ -309,29 +310,52 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in if not user_config_file or not Path(user_config_file).exists(): user_config_file = str(main_folder / "ppp_config.yaml") if user_config_file and Path(user_config_file).exists(): + user_raw: dict[str, Any] = {} with open(user_config_file, "r", encoding="utf-8") as f: - user_raw = yaml.safe_load(f) + user_raw = _yaml_rt.load(f) user_cfg, _ = self.__parse_configuration(user_raw, "user configuration") - else: - user_cfg = None - if user_cfg is not None: - if not isinstance(user_config_file, dict): default_hosts = set((self.config.hosts or {}).keys()) user_hosts = set((user_cfg.hosts or {}).keys()) missing_hosts = default_hosts - user_hosts - if missing_hosts: - self.log( - logging.WARNING, - f"User configuration is missing host(s) from the default: {', '.join(sorted(missing_hosts))}. Consider updating your configuration file.", - ) default_models = set((self.config.models or {}).keys()) user_models = set((user_cfg.models or {}).keys()) missing_models = default_models - user_models - if missing_models: - self.log( - logging.WARNING, - f"User configuration is missing model(s) from the default: {', '.join(sorted(missing_models))}. Consider updating your configuration file.", - ) + if user_raw is not None and (missing_hosts or missing_models): + raw_default_hosts = (default_raw or {}).get("hosts") or {} + raw_default_models = (default_raw or {}).get("models") or {} + if missing_hosts: + self.log( + logging.INFO, + f"Adding missing host(s) from default to user configuration: {', '.join(sorted(missing_hosts))}", + ) + if not user_raw.get("hosts"): + user_raw["hosts"] = {} + for host in sorted(missing_hosts): + if host in raw_default_hosts: + user_raw["hosts"][host] = raw_default_hosts[host] + if missing_models: + self.log( + logging.INFO, + f"Adding missing model(s) from default to user configuration: {', '.join(sorted(missing_models))}", + ) + if not user_raw.get("models"): + user_raw["models"] = {} + for model in sorted(missing_models): + if model in raw_default_models: + user_raw["models"][model] = raw_default_models[model] + try: + with open(user_config_file, "w", encoding="utf-8") as f: + _yaml_rt.dump(user_raw, f) + user_cfg, _ = self.__parse_configuration(user_raw, "user configuration") + self.log( + logging.INFO, + f"Saved updated user configuration to '{escape_single_quotes(user_config_file)}'.", + ) + except Exception as exc: # pylint: disable=broad-exception-caught + self.log(logging.WARNING, f"Failed to save updated user configuration: {exc}") + else: + user_cfg = None + if user_cfg is not None: self.__merge_configuration(user_cfg) self.models_config: dict[str, ModelConfig | None] = self.config.models or {} @@ -455,9 +479,9 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in for model_key, user_model in user_config.models.items(): cfg_model = self.config.models.get(model_key) if user_model is None: - # User wants to disable this model: remove it - if cfg_model is not None: - self.config.models.pop(model_key, None) + # User wants to disable this model: keep the key as None so _is_* variables + # are still set to False (rather than being undefined) + self.config.models[model_key] = None elif cfg_model is None: # New model from user config: add with whatever was specified self.config.models[model_key] = user_model @@ -520,6 +544,9 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in parsed_models: dict[str, ModelConfig | None] = {} raw_models: dict[str, dict | None] = cfg.get("models") or {} for model_key, model_value in raw_models.items(): + if model_value is None: + parsed_models[model_key] = None + continue if not isinstance(model_value, dict) or not any(k in model_value for k in ("detect", "variants")): self.logger.warning( f"{where.capitalize()}: Invalid format for model '{escape_single_quotes(model_key)}'. Discarding model." @@ -674,12 +701,12 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in vs.set_system("_is_pure_" + x, sdchecks[x] and not any(is_models.values())) vs.set_system("_is_variant_" + x, sdchecks[x] and any(is_models.values())) # special cases - vs.set_system("_is_sd", sdchecks["sd1"] or sdchecks["sd2"] or sdchecks["sdxl"] or sdchecks["sd3"]) + vs.set_system("_is_sd", sdchecks.get("sd1", False) or sdchecks.get("sd2", False) or sdchecks.get("sdxl", False) or sdchecks.get("sd3", False)) is_ssd = self.state.env_info.get("is_ssd", False) vs.set_system("_is_ssd", is_ssd) - vs.set_system("_is_sdxl_no_ssd", sdchecks["sdxl"] and not is_ssd) + vs.set_system("_is_sdxl_no_ssd", sdchecks.get("sdxl", False) and not is_ssd) # backcompatibility (but the modern one to use would be _is_pure_sdxl) - vs.set_system("_is_sdxl_no_pony", sdchecks["sdxl"] and not vs.get_system("_is_pony", False)) + vs.set_system("_is_sdxl_no_pony", sdchecks.get("sdxl", False) and not vs.get_system("_is_pony", False)) vs.update_system(input_vars) @@ -1133,8 +1160,12 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in with open(filepath, "a", encoding="utf-8-sig") as f: if not file_exists: f.write("records:\n") + _yaml_dump = _YAML() + _yaml_dump.default_flow_style = False for record in records: - y = yaml.dump(record, allow_unicode=True, default_flow_style=False) + _sio = StringIO() + _yaml_dump.dump(record, _sio) + y = _sio.getvalue() f.write(f" - {textwrap.indent(y, ' ' * 4).strip()}\n") elif ext == ".jsonl": with open(filepath, "a", encoding="utf-8-sig") as f: diff --git a/ppp_common.py b/ppp_common.py index a572603..3639169 100644 --- a/ppp_common.py +++ b/ppp_common.py @@ -7,7 +7,7 @@ import re import textwrap import time import lark -import yaml +from ruamel.yaml import YAML as _YAML from ppp_logging import log from ppp_classes import ONWARNING_CHOICES, PPPInterrupt, PPPState @@ -274,7 +274,7 @@ def convert_a1111_styles_to_wildcard(inp: Path, out: Path): raise RuntimeError(f"No styles found in {inp} to convert to wildcards.") with open(out, "w", encoding="utf-8-sig") as f: f.write(f"# Original names may contain characters that are replaced in the output.\n# Converted from {inp}\n") - yaml.dump(wildcards, f, allow_unicode=True) + _YAML().dump(wildcards, f) def convert_sdnext_styles_to_wildcard(inp: Path, out: Path): @@ -288,7 +288,7 @@ def convert_sdnext_styles_to_wildcard(inp: Path, out: Path): files = inp.glob("*.json") if inp.is_dir() else [inp] for file in files: with open(file, "r", encoding="utf-8-sig") as f: - data = yaml.safe_load(f) + data = _YAML(typ='safe').load(f) if not isinstance(data, list): continue wildcards[file] = {} @@ -307,4 +307,5 @@ def convert_sdnext_styles_to_wildcard(inp: Path, out: Path): f.write("# Original names may contain characters that are replaced in the output.\n") for name, wcs in wildcards.items(): f.write(f"# Converted from {name}\n") - yaml.dump(wcs, f, allow_unicode=True) + yaml_writer = _YAML() + yaml_writer.dump(wcs, f) diff --git a/ppp_enmappings.py b/ppp_enmappings.py index d237644..57d9585 100644 --- a/ppp_enmappings.py +++ b/ppp_enmappings.py @@ -1,7 +1,8 @@ from pathlib import Path from typing import Optional import logging -import yaml +from ruamel.yaml import YAML as _YAML +from ruamel.yaml.error import YAMLError as _YAMLError from ppp_logging import DEBUG_LEVEL, log from ppp_utils import deep_freeze, escape_single_quotes @@ -217,8 +218,8 @@ class PPPExtraNetworkMappings: enmappings_input = enmappings_input.strip() if enmappings_input != "": try: - content = yaml.safe_load(enmappings_input) - except yaml.YAMLError as e: + content = _YAML(typ='safe').load(enmappings_input) + except _YAMLError as e: log( self.__logger, self.__debug_level, @@ -297,7 +298,7 @@ class PPPExtraNetworkMappings: try: try: with open(full_path, "r", encoding="utf-8") as file: - content = yaml.safe_load(file) + content = _YAML(typ='safe').load(file) except: # pylint: disable=bare-except log( self.__logger, @@ -306,7 +307,7 @@ class PPPExtraNetworkMappings: f"Could not read file '{escape_single_quotes(str(full_path))}' with utf-8 encoding, trying windows-1252...", ) with open(full_path, "r", encoding="windows-1252") as file: - content = yaml.safe_load(file) + content = _YAML(typ='safe').load(file) self.__add_extranetwork_mapping(content, full_path) except Exception as e: # pylint: disable=broad-except log( diff --git a/ppp_wildcards.py b/ppp_wildcards.py index 4b5c347..ea8f47c 100644 --- a/ppp_wildcards.py +++ b/ppp_wildcards.py @@ -2,7 +2,8 @@ import fnmatch from pathlib import Path from typing import Any, Optional import logging -import yaml +from ruamel.yaml import YAML as _YAML +from ruamel.yaml.error import YAMLError as _YAMLError from ppp_logging import DEBUG_LEVEL, log from ppp_utils import deep_freeze, escape_single_quotes @@ -217,8 +218,8 @@ class PPPWildcards: wildcards_input = wildcards_input.strip() if wildcards_input != "": try: - content = yaml.safe_load(wildcards_input) - except yaml.YAMLError as e: + content = _YAML(typ='safe').load(wildcards_input) + except _YAMLError as e: log(self.__logger, self.__debug_level, logging.WARNING, f"Invalid format for input wildcards: {e}") return if content is not None: @@ -471,7 +472,7 @@ class PPPWildcards: external_key_parts = list(full_path.with_suffix("").relative_to(base).parts) try: with open(full_path, "r", encoding="utf-8") as file: - content = yaml.safe_load(file) + content = _YAML(typ='safe').load(file) except: # pylint: disable=bare-except log( self.__logger, @@ -480,7 +481,7 @@ class PPPWildcards: f"Could not read file '{escape_single_quotes(str(full_path))}' with utf-8 encoding, trying windows-1252...", ) with open(full_path, "r", encoding="windows-1252") as file: - content = yaml.safe_load(file) + content = _YAML(typ='safe').load(file) self.__add_wildcard(content, full_path, external_key_parts) def __get_wildcards_in_text_file(self, full_path: Path, base: Path): diff --git a/pyproject.toml b/pyproject.toml index 0197849..726873e 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -3,7 +3,7 @@ name = "sd-webui-prompt-postprocessor" description = "Stable Diffusion WebUI & ComfyUI extension to post-process the prompt. Features include: wildcards, sending content from the prompt to the negative prompt, variables, model detection, extranetwork mapping, cleanup." version = "3.2.0" license = { file = "LICENSE.txt" } -dependencies = ["lark", "numpy", "pyyaml", "pydantic"] +dependencies = ["lark", "numpy", "ruamel.yaml", "pydantic"] requires-python = ">=3.10" [project.urls] diff --git a/requirements.txt b/requirements.txt index 27bce30..08b174e 100644 --- a/requirements.txt +++ b/requirements.txt @@ -1,4 +1,4 @@ lark numpy -pyyaml +ruamel.yaml pydantic diff --git a/tests/tests_variants.py b/tests/tests_variants.py index bca077d..2aed154 100644 --- a/tests/tests_variants.py +++ b/tests/tests_variants.py @@ -73,3 +73,33 @@ class TestModelVariants(TestPromptPostProcessorBase): self.extranetwork_maps_obj, ), ) + + def test_variants_null_model(self): + """null model in config disables detection and its variants""" + self.process( + InputTuple( + "SDXLnot SDXL, PONYnot PONY", + "", + ), + OutputTuple("not SDXL, not PONY", ""), + ppp=PromptPostProcessor( + self.ppp_logger, + { + **self.def_env_info, + "model_filename": "./webui/models/Stable-diffusion/ponymodel.safetensors", + "ppp_config": { + "models": { + "sdxl": None, + } + }, + }, + replace( + self.defopts, + on_warning=ONWARNING_CHOICES.warn, + ), + self.grammar_content, + self.interrupt, + self.wildcards_obj, + self.extranetwork_maps_obj, + ), + )