diff --git a/.github/instructions/tests.instructions.md b/.github/instructions/tests.instructions.md index 065245a..f10584e 100644 --- a/.github/instructions/tests.instructions.md +++ b/.github/instructions/tests.instructions.md @@ -97,7 +97,7 @@ Do not use bare `assert` statements. ## Default Options & Environment -Override `self.defopts` or `self.def_env_info` to pass non-default options — do not hardcode option dicts from scratch. +Override `self.defopts` or `self.def_env_info` to pass non-default options. ```python @@ -121,10 +121,10 @@ def test_cl_custom2(self): # only environment changes OutputTuple("a | b", ""), ppp=PromptPostProcessor( self.ppp_logger, - { - **self.def_env_info, - "model_filename": "./webui/models/Stable-diffusion/testmodel.safetensors", - }, + replace( + self.def_env_info, + model_filename="./webui/models/Stable-diffusion/testmodel.safetensors", + ), self.defopts, self.grammar_content, self.interrupt, diff --git a/ppp.py b/ppp.py index 0e0c7a4..3d3a153 100644 --- a/ppp.py +++ b/ppp.py @@ -20,6 +20,7 @@ from ppp_classes import ( HostConfig, ModelConfig, ModelDetectConfig, + PPPEnvInfo, PPPException, RUN_MODE, VariantConfig, @@ -93,7 +94,7 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in def __init__( self, logger: logging.Logger, - env_info: dict[str, Any], + env_info: PPPEnvInfo, options: PPPStateOptions, grammar_content: Optional[str] = None, interrupt: Optional[Callable] = None, @@ -106,7 +107,7 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in Args: logger: The logger object. interrupt: The interrupt function. - env_info: A dictionary with information for the environment and loaded model. + env_info: Environment and model information. options: The options object for configuring PPP behavior. grammar_content: Optional. The grammar content to be used for parsing. wildcards_obj: Optional. The wildcards object to be used for processing wildcards. @@ -283,7 +284,7 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in def log(self, kind, message: str, min_level: DEBUG_LEVEL | None = None, exc_info=None): log(self.logger, self.debug_level, kind, message, min_level, exc_info=exc_info) - def __load_config_and_detect(self, env_info: dict[str, Any]) -> HostConfig: + def __load_config_and_detect(self, env_info: PPPEnvInfo) -> HostConfig: """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") @@ -303,13 +304,13 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in raise PPPInterrupt(errmsg) self.log(logging.WARNING, errmsg) - app = env_info.get("app", "") - user_config_file = env_info.get("ppp_config", "") + app = env_info.app + user_config_file = env_info.ppp_config or "" if isinstance(user_config_file, dict): user_cfg, _ = self.__parse_configuration(user_config_file, "forced configuration") else: if user_config_file == "": - if app == SUPPORTED_APPS.comfyui.value: + if app == SUPPORTED_APPS.comfyui: try: import folder_paths # type: ignore # pylint: disable=import-outside-toplevel,import-error @@ -373,7 +374,7 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in self.known_models: list[str] = list(self.models_config.keys()) # Patch for tests (copy comfyui) - if app == "tests": + if app == SUPPORTED_APPS.tests: if self.config.hosts is None: self.config.hosts = {} self.config.hosts.setdefault("tests", HostConfig()) @@ -384,10 +385,10 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in model.detect = {} model.detect.setdefault("tests", model.detect.get("comfyui", None)) - host_config: HostConfig | None = (self.config.hosts or {}).get(app) + host_config: HostConfig | None = (self.config.hosts or {}).get(app.value) if host_config is None: raise PPPInterrupt( - f"No host configuration found for app '{escape_single_quotes(app)}'. Please check your configuration." + f"No host configuration found for app '{escape_single_quotes(app.value)}'. Please check your configuration." ) # Update env_info with model detection @@ -405,19 +406,19 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in ) self.log( logging.DEBUG, - f"Host configuration ({escape_single_quotes(app)}): {host_config}", + f"Host configuration ({escape_single_quotes(app.value)}): {host_config}", min_level=DEBUG_LEVEL.minimal, ) return host_config - def __run_model_detection(self, env_info: dict[str, Any]) -> None: + def __run_model_detection(self, env_info: PPPEnvInfo) -> None: """Updates the is_* model detection flags in env_info based on model_class.""" - prop_base = env_info.get("property_base", None) - model_class = env_info.get("model_class", "") - model_name = env_info.get("model_filename", "") - app = env_info.get("app", "") - if not model_class and model_name and app == SUPPORTED_APPS.comfyui.value: + prop_base = env_info.property_base + model_class = env_info.model_class + model_name = env_info.model_filename + app = env_info.app + if not model_class and model_name and app == SUPPORTED_APPS.comfyui: try: model_class = get_model_class_from_filename(Path(model_name)) except PPPException as e: @@ -427,26 +428,26 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in min_level=DEBUG_LEVEL.minimal, ) if model_class: - env_info["model_class"] = model_class + env_info.model_class = model_class self.log( logging.DEBUG, f"Detected model class '{model_class}' from filename '{model_name}'", min_level=DEBUG_LEVEL.minimal, ) for m in self.known_models: - env_info["is_" + m] = False + env_info.is_flags[m] = False model_obj = self.models_config.get(m) model_detect = (model_obj.detect if model_obj else None) or {} - model_detect_for_app: ModelDetectConfig | None = model_detect.get(app) + model_detect_for_app: ModelDetectConfig | None = model_detect.get(app.value) if model_detect_for_app is not None: cls_list = model_detect_for_app.class_ or [] if model_class and model_class in cls_list: - env_info["is_" + m] = True + env_info.is_flags[m] = True elif model_detect_for_app.property is not None and prop_base is not None: prop = model_detect_for_app.property attr = getattr(prop_base, prop, None) if isinstance(attr, bool) and attr: - env_info["is_" + m] = True + env_info.is_flags[m] = True def __on_model_info_update(self) -> None: """Called when _modelfullname or _modelclass are set via a prompt command.""" @@ -456,7 +457,7 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in def update( self, - env_info: dict[str, Any], + env_info: PPPEnvInfo, options: PPPStateOptions, wildcards_obj: PPPWildcards, extranetwork_mappings_obj: PPPExtraNetworkMappings, @@ -636,14 +637,21 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in ) @property - def envinfo_hash(self) -> str: - """ - Generates a hash string based on the environment information. - - Returns: - str: A hash string representing the environment information. - """ - return hash(tuple(sorted(self.state.env_info.items()))) + def envinfo_hash(self) -> int: + """Returns a hash of the environment information for cache-invalidation purposes.""" + ei = self.state.env_info + ppp_config_key = ei.ppp_config if isinstance(ei.ppp_config, (str, type(None))) else id(ei.ppp_config) + return hash( + ( + ei.app, + ppp_config_key, + ei.model_class, + ei.model_filename, + ei.models_path, + id(ei.property_base), + tuple(sorted(ei.is_flags.items())), + ) + ) @property def options_hash(self) -> str: @@ -679,19 +687,19 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in vs.set_system(var_name, str(opt_value).split(".", 1)[-1]) # Model related variables - sdchecks = {x: self.state.env_info.get("is_" + x, False) for x in self.known_models} + sdchecks = {x: self.state.env_info.is_flags.get(x, False) for x in self.known_models} # Adding "" as a sentinel that is always True lets next() return "" when no model matches, # giving a well-defined empty-string fallback without a separate None check. sdchecks.update({"": True}) model_name_val = next((k for k, v in sdchecks.items() if v), "") vs.set_system("_model", model_name_val) vs.set_system("_sd", model_name_val) # deprecated - model_filename = self.state.env_info.get("model_filename", "") + model_filename = self.state.env_info.model_filename vs.set_system("_sdfullname", model_filename) # deprecated vs.set_system("_modelfullname", model_filename) vs.set_system("_sdname", Path(model_filename).name) # deprecated vs.set_system("_modelname", Path(model_filename).name) - vs.set_system("_modelclass", self.state.env_info.get("model_class", "")) + vs.set_system("_modelclass", self.state.env_info.model_class) is_models = {} for model_name, model_type_and_substrings in self.variants_definitions.items(): # A variant is only active when its parent model type is currently loaded @@ -726,7 +734,7 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in or sdchecks.get("sdxl", False) or sdchecks.get("sd3", False), ) - is_ssd = self.state.env_info.get("is_ssd", False) + is_ssd = self.state.env_info.is_flags.get("ssd", False) vs.set_system("_is_ssd", 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) @@ -1157,7 +1165,7 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in r"%datetime%": now.strftime(r"%Y-%m-%d_%H-%M-%S"), r"%date%": now.strftime(r"%Y-%m-%d"), r"%time%": now.strftime(r"%H-%M-%S"), - r"%host%": str(self.state.env_info.get("app", "")), + r"%host%": self.state.env_info.app.value, } result = str(self.state.options.results_file) for token, value in substitutions.items(): diff --git a/ppp_classes.py b/ppp_classes.py index f10eb76..ccaaaf8 100644 --- a/ppp_classes.py +++ b/ppp_classes.py @@ -285,12 +285,30 @@ class CyclicalSamplerState: self.last_prompt_pair = None +@dataclass +class PPPEnvInfo: + """Environment and model information passed to PPP at construction time.""" + + app: SUPPORTED_APPS = SUPPORTED_APPS.a1111 + ppp_config: str | dict | None = None + model_class: str = "" + model_filename: str = "" + property_base: Any = None + models_path: str = "" + _is_flags: dict[str, bool] = field(default_factory=dict, init=False, repr=False) + + @property + def is_flags(self) -> dict[str, bool]: + """Boolean model-detection flags keyed by model name (e.g. 'sdxl' -> True).""" + return self._is_flags + + @dataclass(frozen=True) class PPPState: """State object passed to various PPP components during prompt processing.""" logger: Logger - env_info: dict[str, Any] = field(default_factory=dict) + env_info: PPPEnvInfo = field(default_factory=PPPEnvInfo) host_config: HostConfig = field(default_factory=HostConfig) options: PPPStateOptions = field(default_factory=PPPStateOptions) inputs: PPPStateInputs = field(default_factory=PPPStateInputs) diff --git a/ppp_comfyui.py b/ppp_comfyui.py index e9b4fc3..375587c 100644 --- a/ppp_comfyui.py +++ b/ppp_comfyui.py @@ -16,6 +16,7 @@ from ppp_classes import ( DEFAULT_SAMPLER, IFWILDCARDS_CHOICES, ONWARNING_CHOICES, + PPPEnvInfo, SUPPORTED_APPS, PPPException, RUN_MODE, @@ -353,13 +354,13 @@ class PromptPostProcessorComfyUINode: "Modelname was not provided. System model and variant variables will not be properly set.", ) # model class values in ComfyUI\comfy\supported_models.py - env_info = { - "app": SUPPORTED_APPS.comfyui.value, - "models_path": folder_paths.models_dir, - "model_filename": modelname or "", # path is relative to checkpoints folder - "model_class": modelclass, - "property_base": None, - } + env_info = PPPEnvInfo( + app=SUPPORTED_APPS.comfyui, + models_path=folder_paths.models_dir, + model_filename=modelname or "", # path is relative to checkpoints folder + model_class=modelclass, + property_base=None, + ) wildcards_folders = _resolve_wildcards_folders(wc_options["wc_wildcards_folders"] if wc_options else "") enmappings_folders = _resolve_enmappings_folders(en_options["en_mappings_folders"] if en_options else "") @@ -425,7 +426,9 @@ class PromptPostProcessorComfyUINode: ), run_mode=RUN_MODE(run_mode if run_mode else PromptPostProcessor.DEFAULT_RUN_MODE), results_file=results_file, - results_shuffle=rm_options["results_shuffle"] if rm_options else PromptPostProcessor.DEFAULT_RESULTS_SHUFFLE, + results_shuffle=( + rm_options["results_shuffle"] if rm_options else PromptPostProcessor.DEFAULT_RESULTS_SHUFFLE + ), results_limit=rm_options["results_limit"] if rm_options else PromptPostProcessor.DEFAULT_RESULTS_LIMIT, comb_random_fixed=( rm_options["comb_random_fixed"] if rm_options else PromptPostProcessor.DEFAULT_COMB_RANDOM_FIXED diff --git a/ppp_common.py b/ppp_common.py index 2e4e95c..40897ac 100644 --- a/ppp_common.py +++ b/ppp_common.py @@ -265,7 +265,7 @@ def get_model_config_from_filename(filename: Path) -> object | None: config = model_detection.model_config_from_unet(mock_sd, prefix, True) if not config: mock_sd, metadata = comfy.utils.convert_old_quants(mock_sd, "", metadata=None) - #Allow loading unets from checkpoint files + # Allow loading unets from checkpoint files diffusion_model_prefix = model_detection.unet_prefix_from_state_dict(mock_sd) temp_sd = comfy.utils.state_dict_prefix_replace(mock_sd, {diffusion_model_prefix: ""}, filter_keys=True) if len(temp_sd) > 0: @@ -273,9 +273,9 @@ def get_model_config_from_filename(filename: Path) -> object | None: config = model_detection.model_config_from_unet(mock_sd, "", metadata=metadata) if config is None: mock_sd = model_detection.convert_diffusers_mmdit(mock_sd, "") - if mock_sd is not None: #diffusers mmdit + if mock_sd is not None: # diffusers mmdit config = model_detection.model_config_from_unet(mock_sd, "") - else: #diffusers unet + else: # diffusers unet config = model_detection.model_config_from_diffusers_unet(mock_sd) if not config: diff --git a/ppp_enmappings.py b/ppp_enmappings.py index 57d9585..e0b400d 100644 --- a/ppp_enmappings.py +++ b/ppp_enmappings.py @@ -218,7 +218,7 @@ class PPPExtraNetworkMappings: enmappings_input = enmappings_input.strip() if enmappings_input != "": try: - content = _YAML(typ='safe').load(enmappings_input) + content = _YAML(typ="safe").load(enmappings_input) except _YAMLError as e: log( self.__logger, @@ -298,7 +298,7 @@ class PPPExtraNetworkMappings: try: try: with open(full_path, "r", encoding="utf-8") as file: - content = _YAML(typ='safe').load(file) + content = _YAML(typ="safe").load(file) except: # pylint: disable=bare-except log( self.__logger, @@ -307,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(typ='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_tree.py b/ppp_tree.py index f67e789..93997c9 100644 --- a/ppp_tree.py +++ b/ppp_tree.py @@ -1322,14 +1322,14 @@ class TreeProcessor(lark.visitors.Interpreter): f"Invalid variable name '{escape_single_quotes(variable_name)}' detected! System variables cannot be set." ) return - app = self.state.env_info.get("app", "") - if app not in (SUPPORTED_APPS.comfyui.value, SUPPORTED_APPS.tests.value): + app = self.state.env_info.app + if app not in (SUPPORTED_APPS.comfyui, SUPPORTED_APPS.tests): self.warn_or_stop(f"Setting '{escape_single_quotes(variable_name)}' is only supported in ComfyUI.") return evaluated = self.__visit(content, restore_state=False, discard_content=True) - self.state.env_info[settable_sysvars[variable_name]] = evaluated + setattr(self.state.env_info, settable_sysvars[variable_name], evaluated) if variable_name == "_modelfullname": - self.state.env_info["model_class"] = "" # reset model class so it will be re-evaluated + self.state.env_info.model_class = "" # reset so it will be re-evaluated if self.__on_model_info_update is not None: self.__on_model_info_update() info = variable_name + " = " + f"'{escape_single_quotes(evaluated)}'" diff --git a/scripts/ppp_script.py b/scripts/ppp_script.py index d286262..b995710 100644 --- a/scripts/ppp_script.py +++ b/scripts/ppp_script.py @@ -22,6 +22,7 @@ from ppp_classes import ( DEFAULT_SAMPLER, IFWILDCARDS_CHOICES, ONWARNING_CHOICES, + PPPEnvInfo, SUPPORTED_APPS, SUPPORTED_APPS_NAMES, RUN_MODE, @@ -309,7 +310,9 @@ class PromptPostProcessorA1111Script(scripts.Script): results_limit=min(num_seeds, int(input_results_limit)), results_shuffle=input_results_shuffle, comb_random_fixed=input_comb_random_fixed, - default_sampler=DEFAULT_SAMPLER(input_default_sampler if input_default_sampler else PromptPostProcessor.DEFAULT_DEFAULT_SAMPLER) + default_sampler=DEFAULT_SAMPLER( + input_default_sampler if input_default_sampler else PromptPostProcessor.DEFAULT_DEFAULT_SAMPLER + ), ) if not self.ppp_init: self.ppp_init = True @@ -357,17 +360,17 @@ class PromptPostProcessorA1111Script(scripts.Script): logging.INFO, f"Post-processing prompts ({'i2i' if is_i2i else 't2i'})", ) - env_info = { - "app": app.value, - "models_path": models_path, - "model_filename": getattr(p.sd_model.sd_checkpoint_info, "filename", ""), - "model_class": ( + env_info = PPPEnvInfo( + app=app, + models_path=models_path, + model_filename=getattr(p.sd_model.sd_checkpoint_info, "filename", ""), + model_class=( p.sd_model.model_config.__class__.__name__ if app in (SUPPORTED_APPS.forge, SUPPORTED_APPS.forgeneo) else p.sd_model.__class__.__name__ ), - "property_base": p.sd_model, - } + property_base=p.sd_model, + ) wc_wildcards_folders = getattr(opts, "ppp_wil_wildcardsfolders", "") if wc_wildcards_folders == "": wc_wildcards_folders = os.getenv("WILDCARD_DIR", PPPWildcards.DEFAULT_WILDCARDS_FOLDER) @@ -408,7 +411,9 @@ class PromptPostProcessorA1111Script(scripts.Script): self.wildcards_obj, self.extranetwork_mappings_obj, ) - hash_fullenv = hash((self.ppp.envinfo_hash, self.ppp.options_hash, self.wildcards_obj, self.extranetwork_mappings_obj)) + hash_fullenv = hash( + (self.ppp.envinfo_hash, self.ppp.options_hash, self.wildcards_obj, self.extranetwork_mappings_obj) + ) if input_force_equal_seeds: log(self.ppp_logger, self.ppp_debug_level, logging.INFO, "Forcing equal seeds") diff --git a/tests/base_tests.py b/tests/base_tests.py index 83c16dc..51375a6 100644 --- a/tests/base_tests.py +++ b/tests/base_tests.py @@ -1,4 +1,4 @@ -from dataclasses import replace +from dataclasses import replace, make_dataclass import difflib import logging from pathlib import Path @@ -6,7 +6,15 @@ from typing import Any, NamedTuple, Optional import unittest import datetime -from ppp_classes import DEFAULT_SAMPLER, IFWILDCARDS_CHOICES, ONWARNING_CHOICES, RUN_MODE, PPPStateOptions +from ppp_classes import ( + DEFAULT_SAMPLER, + IFWILDCARDS_CHOICES, + ONWARNING_CHOICES, + PPPEnvInfo, + RUN_MODE, + SUPPORTED_APPS, + PPPStateOptions, +) from ppp_enmappings import PPPExtraNetworkMappings # type: ignore from ppp_wildcards import PPPWildcards # type: ignore from ppp import PromptPostProcessor # type: ignore @@ -81,14 +89,14 @@ class TestPromptPostProcessorBase(unittest.TestCase): comb_random_fixed=True, default_sampler=DEFAULT_SAMPLER.random, ) - self.def_env_info = { - "app": "tests", - "ppp_config": None, - "model_class": "SDXL", - "property_base": {"is_sdxl": True}, - "models_path": "./webui/models", - "model_filename": "./webui/models/Stable-diffusion/testmodel.safetensors", - } + self.def_env_info = PPPEnvInfo( + app=SUPPORTED_APPS.tests, + ppp_config=None, + model_class="SDXL", + property_base=make_dataclass("PropertyBase", [("is_sdxl", bool)])(is_sdxl=True), + models_path="./webui/models", + model_filename="./webui/models/Stable-diffusion/testmodel.safetensors", + ) self.interrupted = False self.wildcards_obj = PPPWildcards(self.lf.log) self.extranetwork_maps_obj = PPPExtraNetworkMappings(self.lf.log) diff --git a/tests/tests_host.py b/tests/tests_host.py index 033931d..8ec58b3 100644 --- a/tests/tests_host.py +++ b/tests/tests_host.py @@ -1,3 +1,4 @@ +from dataclasses import replace from ppp import PromptPostProcessor # type: ignore from .base_tests import OutputTuple, InputTuple, TestPromptPostProcessorBase @@ -21,10 +22,10 @@ class TestHosts(TestPromptPostProcessorBase): OutputTuple("(test1:0.9) (test2) (test3:1.5) (test4:0.99)", ""), ppp=PromptPostProcessor( self.ppp_logger, - { - **self.def_env_info, - "ppp_config": {"hosts": {"tests": {"attention": "parentheses"}}}, - }, + replace( + self.def_env_info, + ppp_config={"hosts": {"tests": {"attention": "parentheses"}}}, + ), self.defopts, self.grammar_content, self.interrupt, @@ -42,10 +43,10 @@ class TestHosts(TestPromptPostProcessorBase): OutputTuple("test1 test2 test3", ""), ppp=PromptPostProcessor( self.ppp_logger, - { - **self.def_env_info, - "ppp_config": {"hosts": {"tests": {"attention": "disable"}}}, - }, + replace( + self.def_env_info, + ppp_config={"hosts": {"tests": {"attention": "disable"}}}, + ), self.defopts, self.grammar_content, self.interrupt, @@ -63,10 +64,10 @@ class TestHosts(TestPromptPostProcessorBase): OutputTuple("", ""), ppp=PromptPostProcessor( self.ppp_logger, - { - **self.def_env_info, - "ppp_config": {"hosts": {"tests": {"attention": "remove"}}}, - }, + replace( + self.def_env_info, + ppp_config={"hosts": {"tests": {"attention": "remove"}}}, + ), self.defopts, self.grammar_content, self.interrupt, @@ -84,10 +85,10 @@ class TestHosts(TestPromptPostProcessorBase): OutputTuple("", ""), ppp=PromptPostProcessor( self.ppp_logger, - { - **self.def_env_info, - "ppp_config": {"hosts": {"tests": {"attention": "error"}}}, - }, + replace( + self.def_env_info, + ppp_config={"hosts": {"tests": {"attention": "error"}}}, + ), self.defopts, self.grammar_content, self.interrupt, @@ -106,10 +107,10 @@ class TestHosts(TestPromptPostProcessorBase): OutputTuple("test1", ""), ppp=PromptPostProcessor( self.ppp_logger, - { - **self.def_env_info, - "ppp_config": {"hosts": {"tests": {"scheduling": "before"}}}, - }, + replace( + self.def_env_info, + ppp_config={"hosts": {"tests": {"scheduling": "before"}}}, + ), self.defopts, self.grammar_content, self.interrupt, @@ -127,10 +128,10 @@ class TestHosts(TestPromptPostProcessorBase): OutputTuple("test2", ""), ppp=PromptPostProcessor( self.ppp_logger, - { - **self.def_env_info, - "ppp_config": {"hosts": {"tests": {"scheduling": "after"}}}, - }, + replace( + self.def_env_info, + ppp_config={"hosts": {"tests": {"scheduling": "after"}}}, + ), self.defopts, self.grammar_content, self.interrupt, @@ -148,10 +149,10 @@ class TestHosts(TestPromptPostProcessorBase): OutputTuple("test1 test3", ""), ppp=PromptPostProcessor( self.ppp_logger, - { - **self.def_env_info, - "ppp_config": {"hosts": {"tests": {"scheduling": "first"}}}, - }, + replace( + self.def_env_info, + ppp_config={"hosts": {"tests": {"scheduling": "first"}}}, + ), self.defopts, self.grammar_content, self.interrupt, @@ -169,10 +170,10 @@ class TestHosts(TestPromptPostProcessorBase): OutputTuple("", ""), ppp=PromptPostProcessor( self.ppp_logger, - { - **self.def_env_info, - "ppp_config": {"hosts": {"tests": {"scheduling": "remove"}}}, - }, + replace( + self.def_env_info, + ppp_config={"hosts": {"tests": {"scheduling": "remove"}}}, + ), self.defopts, self.grammar_content, self.interrupt, @@ -190,10 +191,10 @@ class TestHosts(TestPromptPostProcessorBase): OutputTuple("", ""), ppp=PromptPostProcessor( self.ppp_logger, - { - **self.def_env_info, - "ppp_config": {"hosts": {"tests": {"scheduling": "error"}}}, - }, + replace( + self.def_env_info, + ppp_config={"hosts": {"tests": {"scheduling": "error"}}}, + ), self.defopts, self.grammar_content, self.interrupt, @@ -212,10 +213,10 @@ class TestHosts(TestPromptPostProcessorBase): OutputTuple("test1", ""), ppp=PromptPostProcessor( self.ppp_logger, - { - **self.def_env_info, - "ppp_config": {"hosts": {"tests": {"alternation": "first"}}}, - }, + replace( + self.def_env_info, + ppp_config={"hosts": {"tests": {"alternation": "first"}}}, + ), self.defopts, self.grammar_content, self.interrupt, @@ -233,10 +234,10 @@ class TestHosts(TestPromptPostProcessorBase): OutputTuple("", ""), ppp=PromptPostProcessor( self.ppp_logger, - { - **self.def_env_info, - "ppp_config": {"hosts": {"tests": {"alternation": "remove"}}}, - }, + replace( + self.def_env_info, + ppp_config={"hosts": {"tests": {"alternation": "remove"}}}, + ), self.defopts, self.grammar_content, self.interrupt, @@ -254,10 +255,10 @@ class TestHosts(TestPromptPostProcessorBase): OutputTuple("", ""), ppp=PromptPostProcessor( self.ppp_logger, - { - **self.def_env_info, - "ppp_config": {"hosts": {"tests": {"alternation": "error"}}}, - }, + replace( + self.def_env_info, + ppp_config={"hosts": {"tests": {"alternation": "error"}}}, + ), self.defopts, self.grammar_content, self.interrupt, @@ -276,10 +277,10 @@ class TestHosts(TestPromptPostProcessorBase): OutputTuple("test1\ntest2", ""), ppp=PromptPostProcessor( self.ppp_logger, - { - **self.def_env_info, - "ppp_config": {"hosts": {"tests": {"and": "eol"}}}, - }, + replace( + self.def_env_info, + ppp_config={"hosts": {"tests": {"and": "eol"}}}, + ), self.defopts, self.grammar_content, self.interrupt, @@ -297,10 +298,10 @@ class TestHosts(TestPromptPostProcessorBase): OutputTuple("test1, test2", ""), ppp=PromptPostProcessor( self.ppp_logger, - { - **self.def_env_info, - "ppp_config": {"hosts": {"tests": {"and": "comma"}}}, - }, + replace( + self.def_env_info, + ppp_config={"hosts": {"tests": {"and": "comma"}}}, + ), self.defopts, self.grammar_content, self.interrupt, @@ -318,10 +319,10 @@ class TestHosts(TestPromptPostProcessorBase): OutputTuple("test1 test2", ""), ppp=PromptPostProcessor( self.ppp_logger, - { - **self.def_env_info, - "ppp_config": {"hosts": {"tests": {"and": "remove"}}}, - }, + replace( + self.def_env_info, + ppp_config={"hosts": {"tests": {"and": "remove"}}}, + ), self.defopts, self.grammar_content, self.interrupt, @@ -339,10 +340,10 @@ class TestHosts(TestPromptPostProcessorBase): OutputTuple("", ""), ppp=PromptPostProcessor( self.ppp_logger, - { - **self.def_env_info, - "ppp_config": {"hosts": {"tests": {"and": "error"}}}, - }, + replace( + self.def_env_info, + ppp_config={"hosts": {"tests": {"and": "error"}}}, + ), self.defopts, self.grammar_content, self.interrupt, @@ -361,10 +362,10 @@ class TestHosts(TestPromptPostProcessorBase): OutputTuple("test1\ntest2", ""), ppp=PromptPostProcessor( self.ppp_logger, - { - **self.def_env_info, - "ppp_config": {"hosts": {"tests": {"break": "eol"}}}, - }, + replace( + self.def_env_info, + ppp_config={"hosts": {"tests": {"break": "eol"}}}, + ), self.defopts, self.grammar_content, self.interrupt, @@ -382,10 +383,10 @@ class TestHosts(TestPromptPostProcessorBase): OutputTuple("test1, test2", ""), ppp=PromptPostProcessor( self.ppp_logger, - { - **self.def_env_info, - "ppp_config": {"hosts": {"tests": {"break": "comma"}}}, - }, + replace( + self.def_env_info, + ppp_config={"hosts": {"tests": {"break": "comma"}}}, + ), self.defopts, self.grammar_content, self.interrupt, @@ -403,10 +404,10 @@ class TestHosts(TestPromptPostProcessorBase): OutputTuple("test1 test2", ""), ppp=PromptPostProcessor( self.ppp_logger, - { - **self.def_env_info, - "ppp_config": {"hosts": {"tests": {"break": "remove"}}}, - }, + replace( + self.def_env_info, + ppp_config={"hosts": {"tests": {"break": "remove"}}}, + ), self.defopts, self.grammar_content, self.interrupt, @@ -424,10 +425,10 @@ class TestHosts(TestPromptPostProcessorBase): OutputTuple("", ""), ppp=PromptPostProcessor( self.ppp_logger, - { - **self.def_env_info, - "ppp_config": {"hosts": {"tests": {"break": "error"}}}, - }, + replace( + self.def_env_info, + ppp_config={"hosts": {"tests": {"break": "error"}}}, + ), self.defopts, self.grammar_content, self.interrupt, diff --git a/tests/tests_varcomms.py b/tests/tests_varcomms.py index 5f61d39..badd6f4 100644 --- a/tests/tests_varcomms.py +++ b/tests/tests_varcomms.py @@ -1,3 +1,4 @@ +from dataclasses import replace from ppp import PromptPostProcessor # type: ignore from ppp_classes import ONWARNING_CHOICES # type: ignore from .base_tests import OutputTuple, InputTuple, TestPromptPostProcessorBase @@ -770,10 +771,10 @@ class TestVarCommands(TestPromptPostProcessorBase): OutputTuple("this is PONY", ""), ppp=PromptPostProcessor( self.ppp_logger, - { - **self.def_env_info, - "model_filename": "./webui/models/Stable-diffusion/ponymodel.safetensors", - }, + replace( + self.def_env_info, + model_filename="./webui/models/Stable-diffusion/ponymodel.safetensors", + ), self.defopts, self.grammar_content, self.interrupt, @@ -1003,10 +1004,10 @@ class TestVarCommands(TestPromptPostProcessorBase): OutputTuple("inlinetrigger, triggerpony1, triggerpony2", ""), ppp=PromptPostProcessor( self.ppp_logger, - { - **self.def_env_info, - "model_filename": "./webui/models/Stable-diffusion/ponymodel.safetensors", - }, + replace( + self.def_env_info, + model_filename="./webui/models/Stable-diffusion/ponymodel.safetensors", + ), self.defopts, self.grammar_content, self.interrupt, @@ -1024,10 +1025,10 @@ class TestVarCommands(TestPromptPostProcessorBase): OutputTuple("inlinetrigger, triggerpony1, triggerpony2", ""), ppp=PromptPostProcessor( self.ppp_logger, - { - **self.def_env_info, - "model_filename": "./webui/models/Stable-diffusion/ponymodel.safetensors", - }, + replace( + self.def_env_info, + model_filename="./webui/models/Stable-diffusion/ponymodel.safetensors", + ), self.defopts, self.grammar_content, self.interrupt, @@ -1045,10 +1046,10 @@ class TestVarCommands(TestPromptPostProcessorBase): OutputTuple("inlinetrigger, triggerpony1, triggerpony2", ""), ppp=PromptPostProcessor( self.ppp_logger, - { - **self.def_env_info, - "model_filename": "./webui/models/Stable-diffusion/ponymodel.safetensors", - }, + replace( + self.def_env_info, + model_filename="./webui/models/Stable-diffusion/ponymodel.safetensors", + ), self.defopts, self.grammar_content, self.interrupt, @@ -1066,10 +1067,10 @@ class TestVarCommands(TestPromptPostProcessorBase): OutputTuple("inlinetrigger, triggerillustrious1, triggerillustrious2", ""), ppp=PromptPostProcessor( self.ppp_logger, - { - **self.def_env_info, - "model_filename": "./webui/models/Stable-diffusion/ilxlmodel.safetensors", - }, + replace( + self.def_env_info, + model_filename="./webui/models/Stable-diffusion/ilxlmodel.safetensors", + ), self.defopts, self.grammar_content, self.interrupt, diff --git a/tests/tests_variants.py b/tests/tests_variants.py index 2aed154..ca4a65f 100644 --- a/tests/tests_variants.py +++ b/tests/tests_variants.py @@ -24,10 +24,10 @@ class TestModelVariants(TestPromptPostProcessorBase): OutputTuple("test1test2", ""), ppp=PromptPostProcessor( self.ppp_logger, - { - **self.def_env_info, - "model_filename": "./webui/models/Stable-diffusion/testmodel.safetensors", - "ppp_config": { + replace( + self.def_env_info, + model_filename="./webui/models/Stable-diffusion/testmodel.safetensors", + ppp_config={ "models": { "sd1": { "detect": {"tests": {"class": ["SD15", "SD15_instructpix2pix"]}}, @@ -62,7 +62,7 @@ class TestModelVariants(TestPromptPostProcessorBase): }, } }, - }, + ), replace( self.defopts, on_warning=ONWARNING_CHOICES.warn, @@ -84,15 +84,15 @@ class TestModelVariants(TestPromptPostProcessorBase): OutputTuple("not SDXL, not PONY", ""), ppp=PromptPostProcessor( self.ppp_logger, - { - **self.def_env_info, - "model_filename": "./webui/models/Stable-diffusion/ponymodel.safetensors", - "ppp_config": { + replace( + 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,