From b2a5c38f3087729e81ce7b5c420b195ab548778f Mon Sep 17 00:00:00 2001 From: Antonio Cordero Balcazar Date: Thu, 30 Apr 2026 11:32:05 +0200 Subject: [PATCH] * Unify variables in their own object. Co-authored-by: Copilot --- docs/SYNTAX.md | 2 +- ppp.py | 57 +++++++++++------------ ppp_classes.py | 7 ++- ppp_tree.py | 91 ++++++++++++------------------------ ppp_variables.py | 117 +++++++++++++++++++++++++++++++++++++++++++++++ 5 files changed, 178 insertions(+), 96 deletions(-) create mode 100644 ppp_variables.py diff --git a/docs/SYNTAX.md b/docs/SYNTAX.md index dd4173c..c72a7da 100644 --- a/docs/SYNTAX.md +++ b/docs/SYNTAX.md @@ -257,7 +257,7 @@ The star operator `*` is always with inmediate evaluation. And can be accesed/echoed with: * Empty brackets mean the whole array (used when initializing or when echoing the whole array) -* An integer inside the brackets means an indexed value +* An integer inside the brackets means an indexed value. A variable identifier (not an indexed array) can be used to get the integer. * A hash inside the brackets is used to get the length of the array * An ampersand followed by a string inside the brackets (with quotes) is used to get the full array joined with a separator. diff --git a/ppp.py b/ppp.py index f6461d7..60bd8f9 100644 --- a/ppp.py +++ b/ppp.py @@ -22,6 +22,7 @@ from ppp_classes import ( PPPState, PPPStateOptions, ) # pylint: disable=import-error +from ppp_variables import VariableRepository from ppp_logging import DEBUG_LEVEL, log from ppp_tree import TreeProcessor from ppp_utils import escape_single_quotes @@ -228,9 +229,7 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in logger=self.logger, host_config=host_config, options=options, - system_variables={}, - user_variables={}, - echoed_variables={}, + variables=VariableRepository(), wildcards_obj=wildcards_obj, extranetwork_mappings_obj=extranetwork_mappings_obj, parsers={ @@ -547,18 +546,19 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in """ Initializes the system variables. """ - sv = self.state.system_variables - sv.clear() + vs = self.state.variables + vs.clear_system() sdchecks = {x: self.env_info.get("is_" + x, False) for x in self.known_models} sdchecks.update({"": True}) - sv["_model"] = next((k for k, v in sdchecks.items() if v), "") - sv["_sd"] = sv["_model"] # deprecated + 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.env_info.get("model_filename", "") - sv["_sdfullname"] = model_filename # deprecated - sv["_modelfullname"] = model_filename - sv["_sdname"] = os.path.basename(model_filename) # deprecated - sv["_modelname"] = os.path.basename(model_filename) - sv["_modelclass"] = self.env_info.get("model_class", "") + vs.set_system("_sdfullname", model_filename) # deprecated + vs.set_system("_modelfullname", model_filename) + vs.set_system("_sdname", os.path.basename(model_filename)) # deprecated + vs.set_system("_modelname", os.path.basename(model_filename)) + vs.set_system("_modelclass", self.env_info.get("model_class", "")) is_models = {} for model_name, model_type_and_substrings in self.variants_definitions.items(): if not (model_type_and_substrings[0] == "" or sdchecks.get(model_type_and_substrings[0], False)): @@ -574,19 +574,19 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in logging.WARNING, f"Multiple model variants detected at the same time in the filename!: {', '.join(is_models_true)}", ) - sv.update({"_is_" + x: y for x, y in is_models.items()}) + vs.update_system({"_is_" + x: y for x, y in is_models.items()}) for x in sdchecks.keys(): if x != "": - sv["_is_" + x] = sdchecks[x] - sv["_is_pure_" + x] = sdchecks[x] and not any(is_models.values()) - sv["_is_variant_" + x] = sdchecks[x] and any(is_models.values()) + vs.set_system("_is_" + x, sdchecks[x]) + 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 - sv["_is_sd"] = sdchecks["sd1"] or sdchecks["sd2"] or sdchecks["sdxl"] or sdchecks["sd3"] + vs.set_system("_is_sd", sdchecks["sd1"] or sdchecks["sd2"] or sdchecks["sdxl"] or sdchecks["sd3"]) is_ssd = self.env_info.get("is_ssd", False) - sv["_is_ssd"] = is_ssd - sv["_is_sdxl_no_ssd"] = sdchecks["sdxl"] and not is_ssd + vs.set_system("_is_ssd", is_ssd) + vs.set_system("_is_sdxl_no_ssd", sdchecks["sdxl"] and not is_ssd) # backcompatibility (but the modern one to use would be _is_pure_sdxl) - sv["_is_sdxl_no_pony"] = sdchecks["sdxl"] and not sv.get("_is_pony", False) + vs.set_system("_is_sdxl_no_pony", sdchecks["sdxl"] and not vs.get_system("_is_pony", False)) def init_wildcards_options(self): """Initializes the wildcard options.""" @@ -867,9 +867,9 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in Returns: tuple: A tuple containing the processed prompt and negative prompt. """ - self.state.user_variables.clear() - self.state.echoed_variables.clear() - all_variables = {**self.state.system_variables} + self.state.variables.clear_user() + self.state.variables.clear_echoed() + all_variables = self.state.variables.get_all_system() # Process prompt p_processor = TreeProcessor(self.state, rng) @@ -896,15 +896,16 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in negative_prompt = n_processor.start_visit("negative prompt", n_parsed, True) # Complete variables - var_keys = set(self.state.user_variables.keys()).union(set(self.state.echoed_variables.keys())) + var_keys = self.state.variables.all_user_or_echoed_keys() for k in var_keys: - ev = self.state.echoed_variables.get(k) + ev = self.state.variables.get_echoed_value(k) if ev is None: - ev = self.state.user_variables.get(k) + ev = self.state.variables.get_user(k) if ev is None or not isinstance(ev, str): self.log(logging.DEBUG, f"Completing variable: {k}") - ev = p_processor.get_final_user_variable(k) + ev = p_processor.get_final_variable(k) all_variables[k] = self.__cleanup(ev, 0) if self.state.options.cup_cleanup_variables else ev + all_variables = {k: all_variables[k] for k in sorted(all_variables.keys())} self.log(logging.DEBUG, f"All variables: {all_variables}") # Insertions in the negative prompt @@ -1006,7 +1007,7 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in seed = np.random.randint(0, 2**32, dtype=np.int64) prompt = original_prompt negative_prompt = original_negative_prompt - self.log(logging.INFO, f"System variables: {self.state.system_variables}") + self.log(logging.INFO, f"System variables: {self.state.variables.get_all_system()}") self.log(logging.INFO, f"Input seed: {seed}") self.log(logging.INFO, f"Input prompt: {prompt}") self.log(logging.INFO, f"Input negative_prompt: {negative_prompt}") diff --git a/ppp_classes.py b/ppp_classes.py index 29cf90d..1c17b77 100644 --- a/ppp_classes.py +++ b/ppp_classes.py @@ -4,13 +4,14 @@ from dataclasses import dataclass, field from logging import Logger import re from enum import Enum -from typing import Any, Literal, Optional +from typing import Literal, Optional from lark import Lark from pydantic import BaseModel, ConfigDict, Field, field_validator, model_validator from ppp_logging import DEBUG_LEVEL from ppp_wildcards import PPPWildcards from ppp_enmappings import PPPExtraNetworkMappings +from ppp_variables import VariableRepository class SUPPORTED_APPS(Enum): @@ -209,9 +210,7 @@ class PPPState: logger: Logger host_config: HostConfig = field(default_factory=HostConfig) options: PPPStateOptions = field(default_factory=PPPStateOptions) - system_variables: dict[str, Any] = field(default_factory=dict) - user_variables: dict[str, Any] = field(default_factory=dict) - echoed_variables: dict[str, Any] = field(default_factory=dict) + variables: VariableRepository = field(default_factory=VariableRepository) wildcards_obj: PPPWildcards = field(default_factory=PPPWildcards) extranetwork_mappings_obj: PPPExtraNetworkMappings = field(default_factory=PPPExtraNetworkMappings) parsers: dict[str, Lark] = field(default_factory=dict) diff --git a/ppp_tree.py b/ppp_tree.py index e8aa0d2..06a0490 100644 --- a/ppp_tree.py +++ b/ppp_tree.py @@ -102,8 +102,7 @@ class TreeProcessor(lark.visitors.Interpreter): backup_add_at = self.add_at.copy() backup_insertion_at = self.insertion_at.copy() backup_detectedwildcards = self.detectedWildcards.copy() - backup_user_variables = self.state.user_variables.copy() - backup_echoed_variables = self.state.echoed_variables.copy() + backup_vars = self.state.variables.backup_user_and_echoed() if node is not None: if isinstance(node, list): for child in node: @@ -127,10 +126,7 @@ class TreeProcessor(lark.visitors.Interpreter): self.add_at = backup_add_at self.insertion_at = backup_insertion_at self.detectedWildcards = backup_detectedwildcards - self.state.user_variables.clear() - self.state.user_variables.update(backup_user_variables) - self.state.echoed_variables.clear() - self.state.echoed_variables.update(backup_echoed_variables) + self.state.variables.restore_user_and_echoed(backup_vars) return added_result def __get_original_node_content(self, node: lark.Tree | lark.Token, default=None) -> str: @@ -164,26 +160,26 @@ class TreeProcessor(lark.visitors.Interpreter): return None, specifier[2:-1], False if not specifier.isdecimal(): # bare identifier: resolve as variable - specifier = self.get_final_user_variable(specifier) + specifier = self.get_final_variable(specifier) if specifier.isdecimal(): return int(specifier), None, False # invalid specifier return None, None, False - def __get_user_variable_value( + def __get_variable_value( self, name: str, specifier: str | None = None, evaluate=True, visit=False ) -> str | list[str] | None: """ - Get the value of a user variable. + Get the value of a variable. Args: - name (str): The name of the user variable. + name (str): The name of the variable. specifier (str|None): The specifier for an array variable. evaluate (bool): Whether to evaluate the variable. visit (bool): Whether to also visit the variable (add to result). Returns: - str|list[str]|None: The value of the user variable. + str|list[str]|None: The value of the variable. """ def visit_value(v): @@ -200,7 +196,7 @@ class TreeProcessor(lark.visitors.Interpreter): self.result += v return v - v = self.state.user_variables.get(name, None) + v = self.state.variables.get(name) if v is None: return None is_array = name[-2:] == "[]" @@ -250,18 +246,18 @@ class TreeProcessor(lark.visitors.Interpreter): specifier = None return name, specifier - def get_final_user_variable(self, name_specifier: str) -> str: + def get_final_variable(self, name_specifier: str) -> str: """ - Get the final value of a user variable, resolving any references if needed. + Get the final value of a variable, resolving any references if needed. Args: name_specifier (str): The variable reference string. Returns: - str: The final value of the user variable. + str: The final value of the variable. """ name, specifier = self.__separate_arrayref(name_specifier) - v = self.__get_user_variable_value(name, specifier, True, False) + v = self.__get_variable_value(name, specifier, True, False) if isinstance(v, list): _, sep, _ = self.__parse_array_specifier(specifier) if sep is None: @@ -269,26 +265,6 @@ class TreeProcessor(lark.visitors.Interpreter): v = sep.join(str(item) for item in v) return str(v) - def __set_user_variable_value(self, name: str, value: str | lark.Tree | list): - """ - Set the value of a user variable. - - Args: - name (str): The name of the user variable. - value (str|lark.Tree|list): The value to be set. - """ - self.state.user_variables[name] = value - - def __remove_user_variable(self, name: str): - """ - Remove a user variable. - - Args: - name (str): The name of the user variable. - """ - if name in self.state.user_variables: - del self.state.user_variables[name] - def __debug_end(self, construct: str, start_result: str, duration: int, info=None): """ Log the end of a construct processing. @@ -378,15 +354,11 @@ class TreeProcessor(lark.visitors.Interpreter): if c.lower() == "true": return True # Bare identifier - resolve as variable reference - if c.startswith("_"): - vartype = "system" - val = self.state.system_variables.get(c, None) - else: - vartype = "user" - varname, varspecifier = self.__separate_arrayref(c) - val = self.__get_user_variable_value(varname, varspecifier) + varname, varspecifier = self.__separate_arrayref(c) + val = self.__get_variable_value(varname, varspecifier) if val is None: val = "" + vartype = "system" if self.state.variables.name_is_system(c) else "user" self.warn_or_stop(f"Unknown {vartype} variable '{escape_single_quotes(c)}'") if isinstance(val, str): try: @@ -992,7 +964,7 @@ class TreeProcessor(lark.visitors.Interpreter): """ t1 = time.monotonic_ns() start_result = self.result - if variable_name.startswith("_"): + if self.state.variables.name_is_system(variable_name): self.warn_or_stop( f"Invalid variable name '{escape_single_quotes(variable_name)}' detected! System variables cannot be set." ) @@ -1008,7 +980,7 @@ class TreeProcessor(lark.visitors.Interpreter): info = f"{variable_name[0:-2]}[{variable_specifier}]" value_description = self.__get_original_node_content(content, None) value = content - raw_oldvalue = self.state.user_variables.get(variable_name, None) + raw_oldvalue = self.state.variables.get_user(variable_name) newvalue = None some_error = False if variable_specifier is not None: @@ -1083,7 +1055,7 @@ class TreeProcessor(lark.visitors.Interpreter): if vardescriptor_specifier is not None: newvalue = None else: - newvalue = self.__get_user_variable_value(vardescriptor_name, None) + newvalue = self.__get_variable_value(vardescriptor_name, None) elif newvalue.children[0].data == "listvalue": newvalue = list( self.__resolve_operand(c) for c in self.__get_cond_operand(newvalue.children[0]) @@ -1119,8 +1091,8 @@ class TreeProcessor(lark.visitors.Interpreter): newvalue = [newvalue] else: newvalue = raw_oldvalue + [newvalue] - self.__set_user_variable_value(variable_name, newvalue) - currentvalue = self.__get_user_variable_value(variable_name, variable_specifier, False) + self.state.variables.set_user(variable_name, newvalue) + currentvalue = self.__get_variable_value(variable_name, variable_specifier, False) if currentvalue is None: info += "error" elif isinstance(currentvalue, list): @@ -1166,14 +1138,7 @@ class TreeProcessor(lark.visitors.Interpreter): # default_value = self.__visit(default, True) # for log is_array = variable_name[-2:] == "[]" vname = f"{variable_name[0:-2]}[{variable_specifier}]" if variable_specifier is not None else variable_name - if variable_name.startswith("_"): - is_systemvar = True - value = self.state.system_variables.get(variable_name, None) - if value is not None: - self.result += value - else: - is_systemvar = False - value = self.__get_user_variable_value(variable_name, variable_specifier, True, True) + value = self.__get_variable_value(variable_name, variable_specifier, True, True) if value is None: if default is not None: self.log(logging.DEBUG, f"Variable '{escape_single_quotes(vname)}' not found, using default value") @@ -1184,8 +1149,8 @@ class TreeProcessor(lark.visitors.Interpreter): self.warn_or_stop(f"Unknown variable {escape_single_quotes(vname)}") default_value = "" value = "" - if not is_systemvar: - self.state.echoed_variables[vname] = value + if not self.state.variables.name_is_system(variable_name): + self.state.variables.echo(vname, value) t2 = time.monotonic_ns() info = variable_name if is_array and variable_specifier is not None: @@ -1901,9 +1866,9 @@ class TreeProcessor(lark.visitors.Interpreter): vardescriptor_name, vardescriptor_specifier = self.__separate_vardescriptor(var_object.children[0]) variablename = vardescriptor_name # variablevalue = self.__visit(var_object.children[1], False, True) - variablebackup = self.state.user_variables.get(variablename, None) - # self.__remove_user_variable(variablename) - # self.__set_user_variable_value(variablename, variablevalue) + variablebackup = self.state.variables.get_user(variablename) + # self.state.variables.delete_user(variablename) + # self.state.variables.set_user(variablename, variablevalue) self.__varset("wildcard", variablename, vardescriptor_specifier, None, var_object.children[1]) choice_values_all = [] for wildcard in selected_wildcards: @@ -1937,9 +1902,9 @@ class TreeProcessor(lark.visitors.Interpreter): if wildcard_key in self.__wildcard_filters: del self.__wildcard_filters[wildcard_key] if variablename is not None: - self.__remove_user_variable(variablename) + self.state.variables.delete_user(variablename) if variablebackup is not None: - self.state.user_variables[variablename] = variablebackup + self.state.variables.set_user(variablename, variablebackup) elif self.state.options.if_wildcards != IFWILDCARDS_CHOICES.remove: self.detectedWildcards.append(wc) self.result += wc diff --git a/ppp_variables.py b/ppp_variables.py new file mode 100644 index 0000000..9acfb3a --- /dev/null +++ b/ppp_variables.py @@ -0,0 +1,117 @@ +from typing import Any + + +class VariableRepository: + """ + Unified repository for system, user, and echoed prompt variables. + + System variables (underscore-prefixed names like ``_model``) are populated + once per processing session and are read-only during prompt evaluation. + + User variables are created and mutated by set/echo constructs in the prompt. + + Echoed variables record which user variables have already been output and + with what resolved string value. + """ + + def __init__(self) -> None: + self._system: dict[str, Any] = {} + self._user: dict[str, Any] = {} + self._echoed: dict[str, str] = {} + + def name_is_system(self, name: str) -> bool: + """Return True if *name* is a system variable (i.e. starts with an underscore).""" + return name.startswith("_") + + # ---- System variables ---- + + def get_system(self, name: str, default: Any = None) -> Any: + """Return the value of a system variable, or *default* if absent.""" + return self._system.get(name, default) + + def set_system(self, name: str, value: Any) -> None: + """Set a system variable.""" + if not self.name_is_system(name): + raise ValueError(f"invalid system variable name '{name}': must start with an underscore") + self._system[name] = value + + def update_system(self, mapping: dict[str, Any]) -> None: + """Bulk-update system variables from *mapping*.""" + for name in mapping: + if not self.name_is_system(name): + raise ValueError(f"invalid system variable name '{name}': must start with an underscore") + self._system.update(mapping) + + def clear_system(self) -> None: + """Remove all system variables.""" + self._system.clear() + + def get_all_system(self) -> dict[str, Any]: + """Return a shallow copy of all system variables.""" + return {x: self._system[x] for x in sorted(self._system.keys())} + + # ---- User variables ---- + + def get_user(self, name: str, default: Any = None) -> Any: + """Return the value of a user variable, or *default* if absent.""" + return self._user.get(name, default) + + def set_user(self, name: str, value: Any) -> None: + """Set a user variable.""" + if self.name_is_system(name): + raise ValueError(f"invalid user variable name '{name}': must not start with an underscore") + self._user[name] = value + + def delete_user(self, name: str) -> None: + """Remove a user variable (no-op if it does not exist).""" + self._user.pop(name, None) + + def clear_user(self) -> None: + """Remove all user variables.""" + self._user.clear() + + # ---- Echoed variables ---- + + def get_echoed_value(self, name: str, default: str | None = None) -> str | None: + """Return the echoed string value for *name*, or *default* if not echoed.""" + return self._echoed.get(name, default) + + def echo(self, name: str, value: str) -> None: + """Record that *name* was echoed with *value*.""" + self._echoed[name] = value + + def clear_echoed(self) -> None: + """Remove all echoed-variable records.""" + self._echoed.clear() + + # ---- Combined queries ---- + + def get(self, name: str, default: Any = None) -> Any: + """ + Return the value of a variable, checking system variables first, then user variables. + + Args: + name (str): The name of the variable. + default (Any): The value to return if the variable is not found. + """ + if name in self._system: + return self._system.get(name, default) + return self._user.get(name, default) + + def all_user_or_echoed_keys(self) -> set[str]: + """Return the union of user-variable and echoed-variable keys.""" + return set(self._user.keys()) | set(self._echoed.keys()) + + # ---- State backup / restore ---- + + def backup_user_and_echoed(self) -> tuple[dict[str, Any], dict[str, str]]: + """Return shallow-copy snapshots of user and echoed variables for rollback.""" + return self._user.copy(), self._echoed.copy() + + def restore_user_and_echoed(self, backup: tuple[dict[str, Any], dict[str, str]]) -> None: + """Restore user and echoed variables from a snapshot made by :meth:`backup_user_and_echoed`.""" + user_backup, echoed_backup = backup + self._user.clear() + self._user.update(user_backup) + self._echoed.clear() + self._echoed.update(echoed_backup)