* Replaced pyyaml with ruamel.yaml to support updating the user configuration file with minimal fuss.
* Fixed some merging configuration issues.
This commit is contained in:
@@ -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:
|
||||
|
||||
+5
-4
@@ -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)
|
||||
|
||||
+6
-5
@@ -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(
|
||||
|
||||
+6
-5
@@ -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):
|
||||
|
||||
+1
-1
@@ -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]
|
||||
|
||||
+1
-1
@@ -1,4 +1,4 @@
|
||||
lark
|
||||
numpy
|
||||
pyyaml
|
||||
ruamel.yaml
|
||||
pydantic
|
||||
|
||||
@@ -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(
|
||||
"<ppp:if _is_sdxl>SDXL<ppp:else>not SDXL<ppp:/if>, <ppp:if _is_pony>PONY<ppp:else>not PONY<ppp:/if>",
|
||||
"",
|
||||
),
|
||||
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,
|
||||
),
|
||||
)
|
||||
|
||||
Reference in New Issue
Block a user