* 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:
Antonio Cordero Balcazar
2026-05-31 17:40:53 +02:00
parent b0f6335b46
commit 34d0a1f0ae
7 changed files with 105 additions and 41 deletions
+56 -25
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -1,4 +1,4 @@
lark
numpy
pyyaml
ruamel.yaml
pydantic
+30
View File
@@ -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,
),
)