* Unify variables in their own object.

Co-authored-by: Copilot <copilot@github.com>
This commit is contained in:
Antonio Cordero Balcazar
2026-04-30 11:32:05 +02:00
co-authored by Copilot
parent ea93669a82
commit b2a5c38f30
5 changed files with 178 additions and 96 deletions
+1 -1
View File
@@ -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.
+29 -28
View File
@@ -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}")
+3 -4
View File
@@ -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)
+28 -63
View File
@@ -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
+117
View File
@@ -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)