* Unify variables in their own object.
Co-authored-by: Copilot <copilot@github.com>
This commit is contained in:
co-authored by
Copilot
parent
ea93669a82
commit
b2a5c38f30
+1
-1
@@ -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.
|
||||
|
||||
|
||||
@@ -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
@@ -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
@@ -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
|
||||
|
||||
@@ -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)
|
||||
Reference in New Issue
Block a user