* Refactorings.

This commit is contained in:
Antonio Cordero Balcazar
2026-05-16 08:27:31 +02:00
parent 9cb0415e48
commit 1a0ea3c5b0
7 changed files with 24 additions and 24 deletions
+1 -1
View File
@@ -364,7 +364,7 @@ The operands can be a variable (`variable`, `array[]`, `array[index]`), a quoted
When an operand is or contains a variable, it is resolved to the variable's current value before the operation.
Variable values `true` and `false` are considered a boolean, and an all digits value is an integer. Except in substring operations indicated below.
Variable values `true` and `false` are considered a boolean, and numeric content is an integer or float. Except in substring operations indicated below. String comparisons are case insensitive.
The operation can be preceded by `not` for readability, instead of using it in the front.
+8 -5
View File
@@ -285,7 +285,8 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in
def __load_config_and_detect(self, env_info: dict[str, Any]) -> HostConfig:
"""Loads config files, performs model detection, and returns the resolved host config."""
default_config_file = str(Path(__file__).resolve().parent / "ppp_config.yaml.defaults")
main_folder = Path(__file__).resolve().parent
default_config_file = str(main_folder / "ppp_config.yaml.defaults")
try:
with open(default_config_file, "r", encoding="utf-8") as f:
default_raw: dict[str, Any] = yaml.safe_load(f)
@@ -317,7 +318,7 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in
except Exception: # pylint: disable=broad-exception-caught
self.log(logging.WARNING, "Failed to get user directory for PPP config.")
if not user_config_file or not Path(user_config_file).exists():
user_config_file = str(Path(__file__).resolve().parent / "ppp_config.yaml")
user_config_file = str(main_folder / "ppp_config.yaml")
if user_config_file and Path(user_config_file).exists():
with open(user_config_file, "r", encoding="utf-8") as f:
user_raw = yaml.safe_load(f)
@@ -399,7 +400,7 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in
"""Called when _modelfullname or _modelclass are set via a prompt command."""
self.__run_model_detection(self.state.env_info)
self.__init_sysvars()
self.log(logging.INFO, f"Updated system variables: {self.state.variables.get_all_system()}")
self.log(logging.INFO, f"Updated system variables: {self.state.variables.all_system}")
def update(
self,
@@ -576,6 +577,7 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in
result,
)
@property
def envinfo_hash(self) -> str:
"""
Generates a hash string based on the environment information.
@@ -585,6 +587,7 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in
"""
return hash(tuple(sorted(self.state.env_info.items())))
@property
def options_hash(self) -> str:
"""
Generates a hash string based on the options.
@@ -940,7 +943,7 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in
self.log(logging.ERROR, "Interrupting!")
self.interrupt()
v = self.state.variables.get_all_system()
v = self.state.variables.all_system
v.update(variables)
return prompt, negative_prompt, v
@@ -1004,7 +1007,7 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in
def process_prompts_group_start(self):
"""Start of a prompt processing group."""
self.log(logging.INFO, f"System variables: {self.state.variables.get_all_system()}")
self.log(logging.INFO, f"System variables: {self.state.variables.all_system}")
self.log(logging.INFO, f"Combinatorial: {self.state.options.do_combinatorial}")
def process_prompt(
+2 -4
View File
@@ -8,7 +8,7 @@ import nodes # type: ignore
from ppp import PromptPostProcessor
from ppp_classes import IFWILDCARDS_CHOICES, ONWARNING_CHOICES, SUPPORTED_APPS, PPPStateOptions
from ppp_common import get_model_class_from_filename
from ppp_common import get_model_class_from_filename, load_grammar
from ppp_logging import DEBUG_LEVEL, PromptPostProcessorLogFactory, log
from ppp_utils import escape_single_quotes
from ppp_wildcards import PPPWildcards
@@ -68,9 +68,7 @@ class PromptPostProcessorComfyUINode:
def __init__(self):
lf = PromptPostProcessorLogFactory()
self.logger = lf.log
grammar_filename = Path(__file__).resolve().parent / "grammar.lark"
with open(grammar_filename, "r", encoding="utf-8") as file:
self.grammar_content = file.read()
self.grammar_content = load_grammar()
self.wildcards_obj = PPPWildcards(lf.log)
self.extranetwork_mappings_obj = PPPExtraNetworkMappings(lf.log)
self.ppp: PromptPostProcessor | None = None
+1 -1
View File
@@ -172,7 +172,7 @@ class TreeProcessor(lark.visitors.Interpreter):
return results
def __finalize_echoed_variables(self):
var_keys = self.state.variables.all_user_or_echoed_keys()
var_keys = self.state.variables.all_user_or_echoed_keys
for k in var_keys:
ev = self.state.variables.get_echoed_value(k)
if ev is None:
+3 -1
View File
@@ -46,7 +46,8 @@ class VariableRepository:
"""Remove all system variables."""
self._system.clear()
def get_all_system(self) -> dict[str, Any]:
@property
def all_system(self) -> dict[str, Any]:
"""Return a shallow copy of all system variables."""
return self._system.copy()
@@ -98,6 +99,7 @@ class VariableRepository:
return self._system.get(name, default)
return self._user.get(name, default)
@property
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())
+3 -6
View File
@@ -21,6 +21,7 @@ from ppp_logging import DEBUG_LEVEL, PromptPostProcessorLogFactory, log
from ppp_cache import PPPLRUCache
from ppp_wildcards import PPPWildcards
from ppp_enmappings import PPPExtraNetworkMappings
from ppp_common import load_grammar
class PromptPostProcessorA1111Script(scripts.Script):
@@ -64,9 +65,7 @@ class PromptPostProcessorA1111Script(scripts.Script):
super().__init__()
self.instance_index = self.increment_instance_count()
self.name = PromptPostProcessor.NAME
grammar_filename = Path(__file__).resolve().parent.parent / "grammar.lark"
with open(grammar_filename, "r", encoding="utf-8") as file:
self.grammar_content = file.read()
self.grammar_content = load_grammar()
self.ppp_logger = None
self.ppp_debug_level = DEBUG_LEVEL.none.value
self.lru_cache = None
@@ -340,9 +339,7 @@ class PromptPostProcessorA1111Script(scripts.Script):
self.wildcards_obj,
self.extranetwork_mappings_obj,
)
hash_fullenv = hash(
(ppp.envinfo_hash(), ppp.options_hash(), self.wildcards_obj, self.extranetwork_mappings_obj)
)
hash_fullenv = hash((ppp.envinfo_hash, 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")
+6 -6
View File
@@ -11,6 +11,7 @@ from ppp_enmappings import PPPExtraNetworkMappings # type: ignore
from ppp_wildcards import PPPWildcards # type: ignore
from ppp import PromptPostProcessor # type: ignore
from ppp_logging import DEBUG_LEVEL, PromptPostProcessorLogFactory # type: ignore
from ppp_common import load_grammar # type: ignore
class InputTuple(NamedTuple):
@@ -88,11 +89,12 @@ class TestPromptPostProcessorBase(unittest.TestCase):
self.interrupted = False
self.wildcards_obj = PPPWildcards(self.lf.log)
self.extranetwork_maps_obj = PPPExtraNetworkMappings(self.lf.log)
tests_folder = Path(__file__).parent
self.wildcards_obj.refresh_wildcards(
DEBUG_LEVEL.full,
[
Path(__file__).parent / "wildcards",
Path(__file__).parent / "wildcards2",
tests_folder / "wildcards",
tests_folder / "wildcards2",
],
"""
yaml_input:
@@ -104,13 +106,11 @@ class TestPromptPostProcessorBase(unittest.TestCase):
)
self.extranetwork_maps_obj.refresh_extranetwork_mappings(
DEBUG_LEVEL.full,
[Path(__file__).parent / "enmappings"],
[tests_folder / "enmappings"],
"""
""",
)
grammar_filename = Path(__file__).resolve().parent.parent / "grammar.lark"
with open(grammar_filename, "r", encoding="utf-8") as file:
self.grammar_content = file.read()
self.grammar_content = load_grammar()
def interrupt(self):
self.interrupted = True