* Refactorings.
This commit is contained in:
+1
-1
@@ -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.
|
||||
|
||||
|
||||
@@ -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
@@ -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
@@ -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
@@ -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())
|
||||
|
||||
@@ -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
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user