diff --git a/docs/CONFIG.md b/docs/CONFIG.md index 80ebd72..afc1752 100644 --- a/docs/CONFIG.md +++ b/docs/CONFIG.md @@ -58,7 +58,7 @@ Output: ### ACB PPP Wildcards Concat node -This node lets you select up to 10 wildcards that will be concatenated with a chosen separator. +This node lets you select up to 10 wildcards that will be concatenated with a chosen separator. You can't specify wildcard folders in the node, so use the other available options to set them. Inputs: diff --git a/docs/SYNTAX.md b/docs/SYNTAX.md index 5c73226..26a2cdb 100644 --- a/docs/SYNTAX.md +++ b/docs/SYNTAX.md @@ -41,6 +41,7 @@ The construct parameters can be written with the following options (all are opti * "**r**": means it allows repetition of the choices. * "**o**": means it is "optional", and no error will be raised if there are no choices to select from. * "**n**" or "**n-m**" or "**n-**" or "**-m**": number or range of choices to select. Allows zero as the start of a range. Default is 1. +* "**'description'**": optional description, only valid in wildcard definitions. Used only in the Wildcards Concat node in ComfyUI. * "**$$sep**": separator when multiple choices are selected. Default is set in settings. * "**$$**": end of the parameters (not optional if any parameters). @@ -129,11 +130,13 @@ A wildcard definition can be: * A txt file. The wildcard name will be the relative path of the file, without the extension. Each line will be a choice. Lines starting with `#` or empty are ignored. Doesn't support nesting. * An array or scalar value inside a json or yaml file. The wildcard name includes the relative folder path of the file, without the extension, but also the path of the value inside the file (if there is one). If the file contains a dictionary, the filename part is not used for the wildcard name. Supports nesting by having dictionaries inside dictionaries. -The best format is a yaml file with a dictionary of wildcards inside. An editor supporting yaml syntax and linting is recommended (f.e. vscode). +The best format is a yaml file with a dictionary of wildcards inside. An editor supporting yaml syntax and linting is recommended (f.e. VSCode). In a choice, the content after a `#` is ignored. -If the first choice follows the format of wildcard parameters (*including the final `$$`*), it will be used as default parameters for that wildcard (see examples in the tests folder). The choices of the wildcard follow the same format as in the choices construct, or the object format of *Dynamic Prompts* (only in structured files). If using the object format for a choice you can use a new `if` property for the condition, and the `labels` property (an array of strings) and `command` property (a boolean) in addition to the standard `weight` and `text`/`content`. +If the first choice follows the format of wildcard parameters (*including the final `$$`*), it will be used as default parameters for that wildcard (see examples in the tests folder). + +The choices of the wildcard follow the same format as in the choices construct, or the object format of *Dynamic Prompts* (only in structured files). If using the object format for a choice you can use a new `if` property for the condition, and the `labels` property (an array of strings) and `command` property (a boolean) in addition to the standard `weight` and `text`/`content`. ```yaml { command: false, labels: ["some_label"], weight: 2, if: "_is_pony", content: "the text" } # "text" property can be used instead of "content" @@ -142,8 +145,8 @@ If the first choice follows the format of wildcard parameters (*including the fi Wildcard parameters in a json/yaml file can also be in object format, and support two additional properties, prefix and suffix: ```yaml -{ sampler: "~", repeating: false, optional: false, count: 2, prefix: "prefix-", suffix: "-suffix", separator: "/" } -{ sampler: "~", repeating: false, optional: false, from: 2, to: 3, prefix: "prefix-", suffix: "-suffix", separator: "/" } +{ sampler: "~", repeating: false, optional: false, count: 2, description: "test wildcard", prefix: "prefix-", suffix: "-suffix", separator: "/" } +{ sampler: "~", repeating: false, optional: false, from: 2, to: 3, description: "test wildcard", prefix: "prefix-", suffix: "-suffix", separator: "/" } ``` The prefix and suffix are added to the result along with the selected choices and separators. They can contain other constructs, but the separator can't. diff --git a/ppp.py b/ppp.py index 48e271a..5f30edd 100644 --- a/ppp.py +++ b/ppp.py @@ -1,36 +1,20 @@ -import ast +import dataclasses import logging -import math import os import re -import textwrap import time - -from collections import namedtuple -from enum import Enum from typing import Any, Callable, Optional import lark import numpy as np import yaml -from ppp_classes import SUPPORTED_APPS +from ppp_classes import IFWILDCARDS_CHOICES, SUPPORTED_APPS, PPPInterrupt, PPPState, PPPStateOptions from ppp_logging import DEBUG_LEVEL -from ppp_utils import escape_single_quotes -from ppp_wildcards import PPPWildcard, PPPWildcards -from ppp_enmappings import PPPENMappingVariant, PPPExtraNetworkMappings - - -class PPPInterrupt(Exception): - """ - Custom exception to handle interruptions in the PromptPostProcessor. - This exception can be raised to stop the processing of prompts. - """ - - def __init__(self, message: str = "Processing interrupted.", pos_prefix: str = "", neg_prefix: str = ""): - super().__init__(message) - self.message = message - self.pos_prefix = pos_prefix - self.neg_prefix = neg_prefix +from ppp_tree import TreeProcessor +from ppp_utils import escape_single_quotes, format_output +from ppp_common import load_grammar, parse_prompt, preprocess_grammar, warn_or_stop +from ppp_wildcards import PPPWildcards +from ppp_enmappings import PPPExtraNetworkMappings class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-instance-attributes @@ -61,50 +45,40 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in NAME = "Prompt Post-Processor" VERSION = get_version_from_pyproject() - class IFWILDCARDS_CHOICES(Enum): - ignore = "ignore" - remove = "remove" - warn = "warn" - stop = "stop" - - class ONWARNING_CHOICES(Enum): - warn = "warn" - stop = "stop" - - DEFAULT_DEBUG_LEVEL = DEBUG_LEVEL.minimal.value - DEFAULT_ONWARNING = ONWARNING_CHOICES.warn.value - DEFAULT_STN_SEPARATOR = ", " - DEFAULT_STN_IGNORE_REPEATS = True - DEFAULT_WC_PROCESS = True - DEFAULT_IF_WILDCARDS = IFWILDCARDS_CHOICES.stop.value - DEFAULT_CHOICE_SEPARATOR = ", " - DEFAULT_KEEP_CHOICES_ORDER = True - DEFAULT_DO_CLEANUP = True - DEFAULT_CLEANUP_VARIABLES = True - DEFAULT_CUP_EXTRA_SPACES = True - DEFAULT_CUP_EMPTY_CONSTRUCTS = True - DEFAULT_CUP_EXTRA_SEPARATORS = True - DEFAULT_CUP_EXTRA_SEPARATORS2 = True - DEFAULT_CUP_EXTRA_SEPARATORS_INCLUDE_EOL = False - DEFAULT_CUP_BREAKS = False - DEFAULT_CUP_BREAKS_EOL = False - DEFAULT_CUP_ANDS = False - DEFAULT_CUP_ANDS_EOL = False - DEFAULT_CUP_EXTRANETWORK_TAGS = False - DEFAULT_CUP_MERGE_ATTENTION = True - DEFAULT_CUP_REMOVE_EXTRANETWORK_TAGS = False + defopt = {f.name: f.default for f in dataclasses.fields(PPPStateOptions)} + DEFAULT_DEBUG_LEVEL = defopt["debug_level"].value + DEFAULT_ONWARNING = defopt["gen_onwarning"].value + DEFAULT_STN_SEPARATOR = defopt["stn_separator"] + DEFAULT_STN_IGNORE_REPEATS = defopt["stn_ignore_repeats"] + DEFAULT_WC_PROCESS = defopt["wil_process_wildcards"] + DEFAULT_IF_WILDCARDS = defopt["wil_ifwildcards"].value + DEFAULT_CHOICE_SEPARATOR = defopt["wil_choice_separator"] + DEFAULT_KEEP_CHOICES_ORDER = defopt["wil_keep_choices_order"] + DEFAULT_DO_CLEANUP = defopt["cup_do_cleanup"] + DEFAULT_CLEANUP_VARIABLES = defopt["cup_cleanup_variables"] + DEFAULT_CUP_EXTRA_SPACES = defopt["cup_extraspaces"] + DEFAULT_CUP_EMPTY_CONSTRUCTS = defopt["cup_emptyconstructs"] + DEFAULT_CUP_EXTRA_SEPARATORS = defopt["cup_extraseparators"] + DEFAULT_CUP_EXTRA_SEPARATORS2 = defopt["cup_extraseparators2"] + DEFAULT_CUP_EXTRA_SEPARATORS_INCLUDE_EOL = defopt["cup_extraseparators_include_eol"] + DEFAULT_CUP_BREAKS = defopt["cup_breaks"] + DEFAULT_CUP_BREAKS_EOL = defopt["cup_breaks_eol"] + DEFAULT_CUP_ANDS = defopt["cup_ands"] + DEFAULT_CUP_ANDS_EOL = defopt["cup_ands_eol"] + DEFAULT_CUP_EXTRANETWORK_TAGS = defopt["cup_extranetworktags"] + DEFAULT_CUP_MERGE_ATTENTION = defopt["cup_mergeattention"] + DEFAULT_CUP_REMOVE_EXTRANETWORK_TAGS = defopt["rem_removeextranetworktags"] WILDCARD_WARNING = '(WARNING TEXT "INVALID WILDCARD" IN BRIGHT RED:1.5)\nBREAK ' WILDCARD_STOP = "INVALID WILDCARD! {0}\nBREAK " UNPROCESSED_STOP = "UNPROCESSED CONSTRUCTS!\nBREAK " - INVALID_CONTENT_STOP = "INVALID CONTENT! {0}\nBREAK " def __init__( self, logger: logging.Logger, - interrupt: Optional[Callable], env_info: dict[str, Any], - options: Optional[dict[str, Any]] = None, + options: PPPStateOptions, grammar_content: Optional[str] = None, + interrupt: Optional[Callable] = None, wildcards_obj: PPPWildcards = None, extranetwork_mappings_obj: PPPExtraNetworkMappings = None, ): @@ -115,18 +89,14 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in logger: The logger object. interrupt: The interrupt function. env_info: A dictionary with information for the environment and loaded model. - options: Optional. The options dictionary for configuring PPP behavior. + options: The options object for configuring PPP behavior. grammar_content: Optional. The grammar content to be used for parsing. wildcards_obj: Optional. The wildcards object to be used for processing wildcards. extranetwork_mappings_obj: Optional. The extranetwork mappings object to be used for processing. """ self.logger = logger - self.rng = np.random.default_rng() # gets seeded on each process prompt call self.interrupt_callback = interrupt - self.options = options self.env_info = env_info - self.wildcard_obj = wildcards_obj - self.extranetwork_mappings_obj = extranetwork_mappings_obj default_config_file = os.path.join(os.path.dirname(os.path.realpath(__file__)), "ppp_config.yaml.defaults") try: @@ -180,8 +150,8 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in self.models_config[m].get("detect", {}).get("comfyui", None) ) - self.host_config: dict[str, Any] = (self.config.get("hosts") or {}).get(self.env_info.get("app", "")) - if self.host_config is None: + host_config: dict[str, Any] = (self.config.get("hosts") or {}).get(self.env_info.get("app", "")) + if host_config is None: raise PPPInterrupt( f"No host configuration found for app '{escape_single_quotes(self.env_info.get('app', ''))}'. Please check your configuration." ) @@ -201,10 +171,6 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in attr = getattr(prop_base, prop, None) if isinstance(attr, bool) and attr: self.env_info["is_" + m] = True - - # General options - self.debug_level = DEBUG_LEVEL(self.options.get("debug_level", self.DEFAULT_DEBUG_LEVEL)) - self.gen_onwarning = self.ONWARNING_CHOICES(self.options.get("on_warning", self.DEFAULT_ONWARNING)) self.variants_definitions = {} for m in self.known_models: for v, vo in (((self.models_config or {}).get(m) or {}).get("variants") or {}).items(): @@ -214,74 +180,17 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in self.logger.warning( f"Variant name '{escape_single_quotes(v)}' in model '{escape_single_quotes(m)}' conflicts with a known model name. Discarding variant." ) - + self.debug_level = options.debug_level if self.debug_level != DEBUG_LEVEL.none: - self.logger.debug(self.format_output(f"Host configuration: {self.host_config}")) - # Wildcards options - self.wil_process_wildcards = self.options.get("process_wildcards", self.DEFAULT_WC_PROCESS) - self.wil_keep_choices_order = self.options.get("keep_choices_order", self.DEFAULT_KEEP_CHOICES_ORDER) - self.wil_choice_separator = self.options.get("choice_separator", self.DEFAULT_CHOICE_SEPARATOR) - self.wil_ifwildcards = self.IFWILDCARDS_CHOICES(self.options.get("if_wildcards", self.DEFAULT_IF_WILDCARDS)) - # Send to negative options - self.stn_ignore_repeats = self.options.get("stn_ignore_repeats", self.DEFAULT_STN_IGNORE_REPEATS) - self.stn_separator = self.options.get("stn_separator", self.DEFAULT_STN_SEPARATOR) - # Cleanup and remove options - self.cup_do_cleanup = self.options.get("do_cleanup", self.DEFAULT_DO_CLEANUP) - self.cup_cleanup_variables = self.options.get("cleanup_variables", self.DEFAULT_CLEANUP_VARIABLES) - self.cup_extraspaces = self.cup_do_cleanup and self.options.get( - "cleanup_extra_spaces", self.DEFAULT_CUP_EXTRA_SPACES - ) - self.cup_emptyconstructs = self.cup_do_cleanup and self.options.get( - "cleanup_empty_constructs", self.DEFAULT_CUP_EMPTY_CONSTRUCTS - ) - self.cup_extraseparators = self.cup_do_cleanup and self.options.get( - "cleanup_extra_separators", self.DEFAULT_CUP_EXTRA_SEPARATORS - ) - self.cup_extraseparators2 = self.cup_do_cleanup and self.options.get( - "cleanup_extra_separators2", self.DEFAULT_CUP_EXTRA_SEPARATORS2 - ) - self.cup_extraseparators_include_eol = self.cup_do_cleanup and self.options.get( - "cleanup_extra_separators_include_eol", self.DEFAULT_CUP_EXTRA_SEPARATORS_INCLUDE_EOL - ) - self.cup_breaks = self.cup_do_cleanup and self.options.get("cleanup_breaks", self.DEFAULT_CUP_BREAKS) - self.cup_breaks_eol = self.cup_do_cleanup and self.options.get( - "cleanup_breaks_eol", self.DEFAULT_CUP_BREAKS_EOL - ) - self.cup_ands = self.cup_do_cleanup and self.options.get("cleanup_ands", self.DEFAULT_CUP_ANDS) - self.cup_ands_eol = self.cup_do_cleanup and self.options.get("cleanup_ands_eol", self.DEFAULT_CUP_ANDS_EOL) - self.cup_extranetworktags = self.cup_do_cleanup and self.options.get( - "cleanup_extranetwork_tags", self.DEFAULT_CUP_EXTRANETWORK_TAGS - ) - self.cup_mergeattention = self.cup_do_cleanup and self.options.get( - "cleanup_merge_attention", self.DEFAULT_CUP_MERGE_ATTENTION - ) - self.rem_removeextranetworktags = self.cup_do_cleanup and self.options.get( - "remove_extranetwork_tags", self.DEFAULT_CUP_REMOVE_EXTRANETWORK_TAGS - ) + self.logger.debug(format_output(f"Host configuration: {host_config}")) # if self.debug_level != DEBUG_LEVEL.none: # self.logger.info(f"Detected environment info: {env_info}") - # Process with lark (debug with https://www.lark-parser.org/ide/) if grammar_content is None: - grammar_filename = os.path.join(os.path.dirname(os.path.realpath(__file__)), "grammar.lark") - with open(grammar_filename, "r", encoding="utf-8") as file: - grammar_content = file.read() - + grammar_content = load_grammar() # Preprocess grammar content for conditional compilation - self.parser_full_only_old = lark.Lark( - self.__preprocess_grammar( - grammar_content, - { - "ALLOW_NEW_CONTENT": False, - "ALLOW_WILDCARDS": False, - "ALLOW_CHOICES": False, - "ALLOW_COMMVARS": False, - }, - ), - propagate_positions=True, - ) - grammar_content_full = self.__preprocess_grammar( + grammar_content_full = preprocess_grammar( grammar_content, { "ALLOW_NEW_CONTENT": True, @@ -290,112 +199,134 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in "ALLOW_COMMVARS": True, }, ) - self.parser_complete_full = lark.Lark( - grammar_content_full, - propagate_positions=True, - ) - self.parser_complete_wc_ch = lark.Lark( - self.__preprocess_grammar( - grammar_content, - { - "ALLOW_NEW_CONTENT": True, - "ALLOW_WILDCARDS": True, - "ALLOW_CHOICES": True, - "ALLOW_COMMVARS": False, - }, - ), - propagate_positions=True, - ) - self.parser_complete_wc_cv = lark.Lark( - self.__preprocess_grammar( - grammar_content, - { - "ALLOW_NEW_CONTENT": True, - "ALLOW_WILDCARDS": True, - "ALLOW_CHOICES": False, - "ALLOW_COMMVARS": True, - }, - ), - propagate_positions=True, - ) - self.parser_complete_ch_cv = lark.Lark( - self.__preprocess_grammar( - grammar_content, - { - "ALLOW_NEW_CONTENT": True, - "ALLOW_WILDCARDS": False, - "ALLOW_CHOICES": True, - "ALLOW_COMMVARS": True, - }, - ), - propagate_positions=True, - ) - self.parser_complete_wc = lark.Lark( - self.__preprocess_grammar( - grammar_content, - { - "ALLOW_NEW_CONTENT": True, - "ALLOW_WILDCARDS": True, - "ALLOW_CHOICES": False, - "ALLOW_COMMVARS": False, - }, - ), - propagate_positions=True, - ) - self.parser_complete_ch = lark.Lark( - self.__preprocess_grammar( - grammar_content, - { - "ALLOW_NEW_CONTENT": True, - "ALLOW_WILDCARDS": False, - "ALLOW_CHOICES": True, - "ALLOW_COMMVARS": False, - }, - ), - propagate_positions=True, - ) - self.parser_complete_cv = lark.Lark( - self.__preprocess_grammar( - grammar_content, - { - "ALLOW_NEW_CONTENT": True, - "ALLOW_WILDCARDS": False, - "ALLOW_CHOICES": False, - "ALLOW_COMMVARS": True, - }, - ), - propagate_positions=True, - ) - # Partial parsers - self.parser_content = lark.Lark( - grammar_content_full, - propagate_positions=True, - start="content", - ) - self.parser_choice = lark.Lark( - grammar_content_full, - propagate_positions=True, - start="choice", - ) - self.parser_wcdefoptions = lark.Lark( - grammar_content_full, - propagate_positions=True, - start="wcdefoptions", - ) - self.parser_condition = lark.Lark( - grammar_content_full, - propagate_positions=True, - start="condition", - ) - self.parser_choicevalue = lark.Lark( - grammar_content_full, - propagate_positions=True, - start="choicevalue", + self.state = PPPState( + logger=self.logger, + host_config=host_config, + options=options, + system_variables={}, + user_variables={}, + echoed_variables={}, + wildcards_obj=wildcards_obj, + extranetwork_mappings_obj=extranetwork_mappings_obj, + parsers={ + "full": lark.Lark( + grammar_content_full, + propagate_positions=True, + ), + "wc_ch": lark.Lark( + preprocess_grammar( + grammar_content, + { + "ALLOW_NEW_CONTENT": True, + "ALLOW_WILDCARDS": True, + "ALLOW_CHOICES": True, + "ALLOW_COMMVARS": False, + }, + ), + propagate_positions=True, + ), + "wc_cv": lark.Lark( + preprocess_grammar( + grammar_content, + { + "ALLOW_NEW_CONTENT": True, + "ALLOW_WILDCARDS": True, + "ALLOW_CHOICES": False, + "ALLOW_COMMVARS": True, + }, + ), + propagate_positions=True, + ), + "ch_cv": lark.Lark( + preprocess_grammar( + grammar_content, + { + "ALLOW_NEW_CONTENT": True, + "ALLOW_WILDCARDS": False, + "ALLOW_CHOICES": True, + "ALLOW_COMMVARS": True, + }, + ), + propagate_positions=True, + ), + "wc": lark.Lark( + preprocess_grammar( + grammar_content, + { + "ALLOW_NEW_CONTENT": True, + "ALLOW_WILDCARDS": True, + "ALLOW_CHOICES": False, + "ALLOW_COMMVARS": False, + }, + ), + propagate_positions=True, + ), + "ch": lark.Lark( + preprocess_grammar( + grammar_content, + { + "ALLOW_NEW_CONTENT": True, + "ALLOW_WILDCARDS": False, + "ALLOW_CHOICES": True, + "ALLOW_COMMVARS": False, + }, + ), + propagate_positions=True, + ), + "cv": lark.Lark( + preprocess_grammar( + grammar_content, + { + "ALLOW_NEW_CONTENT": True, + "ALLOW_WILDCARDS": False, + "ALLOW_CHOICES": False, + "ALLOW_COMMVARS": True, + }, + ), + propagate_positions=True, + ), + "only_old": lark.Lark( + preprocess_grammar( + grammar_content, + { + "ALLOW_NEW_CONTENT": False, + "ALLOW_WILDCARDS": False, + "ALLOW_CHOICES": False, + "ALLOW_COMMVARS": False, + }, + ), + propagate_positions=True, + ), + # Partial parsers + "content": lark.Lark( + grammar_content_full, + propagate_positions=True, + start="content", + ), + "choice": lark.Lark( + grammar_content_full, + propagate_positions=True, + start="choice", + ), + "wcdefoptions": lark.Lark( + grammar_content_full, + propagate_positions=True, + start="wcdefoptions", + ), + "condition": lark.Lark( + grammar_content_full, + propagate_positions=True, + start="condition", + ), + "choicevalue": lark.Lark( + grammar_content_full, + propagate_positions=True, + start="choicevalue", + ), + }, ) self.__init_sysvars() - self.user_variables = {} - self.echoed_variables = {} def __merge_configuration(self, user_config): """ @@ -623,131 +554,28 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in Returns: str: A hash string representing the options. """ - return hash(tuple(sorted(self.options.items()))) - - def __preprocess_grammar(self, grammar_content: str, options: dict[str, bool]) -> str: - """ - Preprocesses the grammar content to handle conditional compilation directives. - - Args: - grammar_content (str): The raw grammar content. - options (dict[str,bool]): Options for preprocessing. - - Returns: - str: The preprocessed grammar content. - """ - lines = grammar_content.split("\n") - result_lines = [] - skip_current_block = [] - all_blocks_skipped = [] - - def eval_bool_expr(expr: str, constants: dict[str, bool]) -> bool: - """ - Evaluates a boolean expression using known constants. - Supports: and, or, not, parentheses, and named constants. - - Args: - expr (str): The boolean expression to evaluate. - constants (dict[str, bool]): A dictionary of constant values. - Returns: - bool: The result of the evaluated expression. - """ - tree = ast.parse(expr, mode="eval") - - def _eval(node) -> bool: - if isinstance(node, ast.Expression): - return _eval(node.body) - if isinstance(node, ast.BoolOp): - if isinstance(node.op, ast.And): - return all(_eval(v) for v in node.values) - if isinstance(node.op, ast.Or): - return any(_eval(v) for v in node.values) - if isinstance(node, ast.UnaryOp) and isinstance(node.op, ast.Not): - return not _eval(node.operand) - if isinstance(node, ast.Name): - return bool(constants[node.id]) # raises KeyError for unknown names - if isinstance(node, ast.Constant) and isinstance(node.value, bool): - return node.value - raise ValueError(f"Unsupported construct: {ast.dump(node)}") - - return _eval(tree) - - for line in lines: - stripped_line = line.strip() - if stripped_line.startswith("//#if"): - # Extract condition from the #if directive - conditions = stripped_line[5:].strip() - # Evaluate the conditions - skip_current_block.append(not eval_bool_expr(conditions, options)) - all_blocks_skipped.append(skip_current_block[-1]) - elif stripped_line.startswith("//#elif"): - if not skip_current_block: - self.logger.warning("Unmatched //#elif directive found in grammar content.") - elif all_blocks_skipped[-1]: - # Extract condition from the #elif directive - conditions = stripped_line[7:].strip() - # Evaluate the conditions - skip_current_block[-1] = not eval_bool_expr(conditions, options) - if not skip_current_block[-1]: - all_blocks_skipped[-1] = False - else: - skip_current_block[-1] = True - elif stripped_line.startswith("//#else"): - if not skip_current_block: - self.logger.warning("Unmatched //#else directive found in grammar content.") - elif all_blocks_skipped[-1]: - skip_current_block[-1] = False - else: - skip_current_block[-1] = True - elif stripped_line.startswith("//#endif"): - if not skip_current_block: - self.logger.warning("Unmatched //#endif directive found in grammar content.") - else: - skip_current_block.pop() - all_blocks_skipped.pop() - elif stripped_line.startswith("//#"): - self.logger.warning(f"Unrecognized directive found in grammar content: {stripped_line}") - elif not any(skip_current_block): - # Include the line if we're not skipping any current block - result_lines.append(stripped_line) - # Check for unclosed blocks at the end - if skip_current_block: - raise PPPInterrupt( - f"Found {len(skip_current_block)} unclosed conditional directive(s) at the end of the grammar file" - ) - return "\n".join(result_lines) + return hash(self.state.options) def interrupt(self): if self.interrupt_callback is not None: self.interrupt_callback() - def format_output(self, text: str) -> str: - """ - Formats the output text by encoding it using unicode_escape and decoding it using utf-8. - - Args: - text (str): The input text to be formatted. - - Returns: - str: The formatted output text. - """ - return text.encode("unicode_escape").decode("utf-8") - def __init_sysvars(self): """ Initializes the system variables. """ - self.system_variables = {} + sv = self.state.system_variables + sv.clear() sdchecks = {x: self.env_info.get("is_" + x, False) for x in self.known_models} sdchecks.update({"": True}) - self.system_variables["_model"] = next((k for k, v in sdchecks.items() if v), "") - self.system_variables["_sd"] = self.system_variables["_model"] # deprecated + sv["_model"] = next((k for k, v in sdchecks.items() if v), "") + sv["_sd"] = sv["_model"] # deprecated model_filename = self.env_info.get("model_filename", "") - self.system_variables["_sdfullname"] = model_filename # deprecated - self.system_variables["_modelfullname"] = model_filename - self.system_variables["_sdname"] = os.path.basename(model_filename) # deprecated - self.system_variables["_modelname"] = os.path.basename(model_filename) - self.system_variables["_modelclass"] = self.env_info.get("model_class", "") + 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", "") 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)): @@ -762,21 +590,25 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in self.logger.warning( f"Multiple model variants detected at the same time in the filename!: {', '.join(is_models_true)}" ) - self.system_variables.update({"_is_" + x: y for x, y in is_models.items()}) + sv.update({"_is_" + x: y for x, y in is_models.items()}) for x in sdchecks.keys(): if x != "": - self.system_variables["_is_" + x] = sdchecks[x] - self.system_variables["_is_pure_" + x] = sdchecks[x] and not any(is_models.values()) - self.system_variables["_is_variant_" + x] = sdchecks[x] and any(is_models.values()) + 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()) # special cases - self.system_variables["_is_sd"] = sdchecks["sd1"] or sdchecks["sd2"] or sdchecks["sdxl"] or sdchecks["sd3"] + sv["_is_sd"] = sdchecks["sd1"] or sdchecks["sd2"] or sdchecks["sdxl"] or sdchecks["sd3"] is_ssd = self.env_info.get("is_ssd", False) - self.system_variables["_is_ssd"] = is_ssd - self.system_variables["_is_sdxl_no_ssd"] = sdchecks["sdxl"] and not is_ssd + sv["_is_ssd"] = is_ssd + sv["_is_sdxl_no_ssd"] = sdchecks["sdxl"] and not is_ssd # backcompatibility (but the modern one to use would be _is_pure_sdxl) - self.system_variables["_is_sdxl_no_pony"] = sdchecks["sdxl"] and not self.system_variables.get( - "_is_pony", False - ) + sv["_is_sdxl_no_pony"] = sdchecks["sdxl"] and not sv.get("_is_pony", False) + + def init_wildcards_options(self): + """Initializes the wildcard options.""" + _tree = TreeProcessor(self.state, np.random.default_rng()) + for wc in self.state.wildcards_obj.wildcards.values(): + _tree.get_wildcard_options(wc) def __add_to_insertion_points( self, negative_prompt: str, add_at_insertion_point: list[str], insertion_at: list[tuple[int, int]] @@ -799,22 +631,28 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in if insertion_at[n] is not None: ipp = insertion_at[n][0] ipl = insertion_at[n][1] - insertion_at[n][0] - if negative_prompt[ipp - len(self.stn_separator) : ipp] == self.stn_separator: - ipp -= len(self.stn_separator) # adjust for existing start separator - ipl += len(self.stn_separator) + if ( + negative_prompt[ipp - len(self.state.options.stn_separator) : ipp] + == self.state.options.stn_separator + ): + ipp -= len(self.state.options.stn_separator) # adjust for existing start separator + ipl += len(self.state.options.stn_separator) add_at_insertion_point[n].insert(0, negative_prompt[:ipp]) - if negative_prompt[ipp + ipl : ipp + ipl + len(self.stn_separator)] == self.stn_separator: - ipl += len(self.stn_separator) # adjust for existing end separator + if ( + negative_prompt[ipp + ipl : ipp + ipl + len(self.state.options.stn_separator)] + == self.state.options.stn_separator + ): + ipl += len(self.state.options.stn_separator) # adjust for existing end separator endPart = negative_prompt[ipp + ipl :] if len(endPart) > 0: add_at_insertion_point[n].append(endPart) - negative_prompt = self.stn_separator.join(add_at_insertion_point[n]) + negative_prompt = self.state.options.stn_separator.join(add_at_insertion_point[n]) else: ipp = 0 - if negative_prompt.startswith(self.stn_separator): - ipp = len(self.stn_separator) + if negative_prompt.startswith(self.state.options.stn_separator): + ipp = len(self.state.options.stn_separator) add_at_insertion_point[n].append(negative_prompt[ipp:]) - negative_prompt = self.stn_separator.join(add_at_insertion_point[n]) + negative_prompt = self.state.options.stn_separator.join(add_at_insertion_point[n]) return negative_prompt def __add_to_start(self, negative_prompt: str, add_at_start: list[str]) -> str: @@ -830,10 +668,10 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in """ if len(negative_prompt) > 0: ipp = 0 - if negative_prompt.startswith(self.stn_separator): - ipp = len(self.stn_separator) # adjust for existing end separator + if negative_prompt.startswith(self.state.options.stn_separator): + ipp = len(self.state.options.stn_separator) # adjust for existing end separator add_at_start.append(negative_prompt[ipp:]) - negative_prompt = self.stn_separator.join(add_at_start) + negative_prompt = self.state.options.stn_separator.join(add_at_start) return negative_prompt def __add_to_end(self, negative_prompt: str, add_at_end: list[str]) -> str: @@ -849,10 +687,10 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in """ if len(negative_prompt) > 0: ipl = len(negative_prompt) - if negative_prompt.endswith(self.stn_separator): - ipl -= len(self.stn_separator) # adjust for existing start separator + if negative_prompt.endswith(self.state.options.stn_separator): + ipl -= len(self.state.options.stn_separator) # adjust for existing start separator add_at_end.insert(0, negative_prompt[:ipl]) - negative_prompt = self.stn_separator.join(add_at_end) + negative_prompt = self.state.options.stn_separator.join(add_at_end) return negative_prompt def __cleanup(self, text: str, where: int = 0) -> str: @@ -866,12 +704,12 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in Returns: str: The resulting text. """ - break_processing = self.host_config.get("break", "ok") + break_processing = self.state.host_config.get("break", "ok") # break_processing == "ok" (and always) - if self.cup_breaks_eol: + if self.state.options.cup_breaks_eol: # replace spaces before break with EOL text = re.sub(r"[, ]+BREAK\b", "\nBREAK", text) - if self.cup_breaks: + if self.state.options.cup_breaks: # collapse separators and commas before BREAK text = re.sub(r"[, ]+BREAK\b", " BREAK", text) # collapse separators and commas after BREAK @@ -901,9 +739,9 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in self.logger.debug(f"BREAK construct {break_replacements[break_processing][0]}") elif break_processing == "error": if re.search(r"\bBREAK\b", text): - self.warn_or_stop(where == -1, "BREAK constructs are not allowed!") + warn_or_stop(self.state, where == -1, "BREAK constructs are not allowed!") - if self.cup_ands: + if self.state.options.cup_ands: # collapse ANDs with space after text = re.sub(r"\bAND(?:\s+AND)+\s+", "AND ", text) # collapse ANDs without space after @@ -917,15 +755,15 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in # remove at end of prompt text = re.sub(r"(\s*\bAND)+\Z", "", text) - escapedSeparator = re.escape(self.stn_separator) - optwhitespace = r"\s*" if self.cup_extraseparators_include_eol else r"[ \t\v\f]*" + escapedSeparator = re.escape(self.state.options.stn_separator) + optwhitespace = r"\s*" if self.state.options.cup_extraseparators_include_eol else r"[ \t\v\f]*" optwhitespace_separator = optwhitespace + escapedSeparator + optwhitespace optwhitespace_comma = optwhitespace + "," + optwhitespace - sep_options = [(optwhitespace_separator, self.stn_separator)] # sendtonegative separator + sep_options = [(optwhitespace_separator, self.state.options.stn_separator)] # sendtonegative separator if optwhitespace_comma != optwhitespace_separator: sep_options.append((optwhitespace_comma, ", ")) # regular comma separator for sep, replacement in sep_options: - if self.cup_extraseparators: + if self.state.options.cup_extraseparators: # collapse separators text = re.sub(r"(?:" + sep + r"){2,}", replacement, text) # remove separator after starting parenthesis, starting bracket @@ -948,17 +786,17 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in r"\1", text, ) - if self.cup_extraseparators2: + if self.state.options.cup_extraseparators2: # remove at start of prompt or line text = re.sub(r"^(?:" + sep + r")+", "", text, flags=re.MULTILINE) # remove at end of prompt or line text = re.sub(r"(?:" + sep + r")+$", "", text, flags=re.MULTILINE) - if self.cup_extranetworktags: + if self.state.options.cup_extranetworktags: # remove spaces before < text = re.sub(r"\B\s+<(?!!)", "<", text) # remove spaces after > text = re.sub(r">\s+\B", ">", text) - if self.cup_extraspaces: + if self.state.options.cup_extraspaces: # remove spaces before comma text = re.sub(r"[ ]+,", ",", text) # remove spaces at end of line @@ -996,65 +834,67 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in } if tests["ALLOW_WILDCARDS"] and tests["ALLOW_CHOICES"] and tests["ALLOW_COMMVARS"]: return ( - self.parser_complete_full, + self.state.parsers["full"], "full parser with wildcards, choices, commands and variables", ) if tests["ALLOW_WILDCARDS"] and tests["ALLOW_CHOICES"]: return ( - self.parser_complete_wc_ch, + self.state.parsers["wc_ch"], "parser with wildcards and choices", ) if tests["ALLOW_WILDCARDS"] and tests["ALLOW_COMMVARS"]: return ( - self.parser_complete_wc_cv, + self.state.parsers["wc_cv"], "parser with wildcards, commands and variables", ) if tests["ALLOW_CHOICES"] and tests["ALLOW_COMMVARS"]: return ( - self.parser_complete_ch_cv, + self.state.parsers["ch_cv"], "parser with choices, commands and variables", ) if tests["ALLOW_WILDCARDS"]: return ( - self.parser_complete_wc, + self.state.parsers["wc"], "parser with wildcards", ) if tests["ALLOW_CHOICES"]: return ( - self.parser_complete_ch, + self.state.parsers["ch"], "parser with choices", ) if tests["ALLOW_COMMVARS"]: return ( - self.parser_complete_cv, + self.state.parsers["cv"], "parser with commands and variables", ) return ( - self.parser_full_only_old, + self.state.parsers["only_old"], "simple parser without new constructs", ) - def __processprompts(self, prompt, negative_prompt): + def __processprompts(self, rng, prompt, negative_prompt): """ Process the prompt and negative prompt. Args: + rng (numpy.random.Generator): The random number generator. prompt (str): The prompt. negative_prompt (str): The negative prompt. Returns: tuple: A tuple containing the processed prompt and negative prompt. """ - self.user_variables = {} - self.echoed_variables = {} - all_variables = {**self.system_variables} + self.state.user_variables.clear() + self.state.echoed_variables.clear() + all_variables = {**self.state.system_variables} # Process prompt - p_processor = self.TreeProcessor(self) + p_processor = TreeProcessor(self.state, rng) (prompt_parser, parser_description) = self.__get_best_parser(prompt) if self.debug_level == DEBUG_LEVEL.full: self.logger.debug(f"Using {parser_description} for prompt") - p_parsed = self.parse_prompt( + p_parsed = parse_prompt( + self.state, "prompt", prompt, prompt_parser, @@ -1062,11 +902,12 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in prompt = p_processor.start_visit("prompt", p_parsed, False) # Process negative prompt - n_processor = self.TreeProcessor(self) + n_processor = TreeProcessor(self.state, rng) (n_prompt_parser, n_parser_description) = self.__get_best_parser(negative_prompt) if self.debug_level == DEBUG_LEVEL.full: self.logger.debug(f"Using {n_parser_description} for negative prompt") - n_parsed = self.parse_prompt( + n_parsed = parse_prompt( + self.state, "negative prompt", negative_prompt, n_prompt_parser, @@ -1074,23 +915,23 @@ 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.user_variables.keys()).union(set(self.echoed_variables.keys())) + var_keys = set(self.state.user_variables.keys()).union(set(self.state.echoed_variables.keys())) for k in var_keys: - ev = self.echoed_variables.get(k) + ev = self.state.echoed_variables.get(k) if ev is None: - ev = self.user_variables.get(k) + ev = self.state.user_variables.get(k) if ev is None or not isinstance(ev, str): if self.debug_level == DEBUG_LEVEL.full: - self.logger.debug(self.format_output(f"Completing variable: {k}")) + self.logger.debug(format_output(f"Completing variable: {k}")) ev = p_processor.get_final_user_variable(k) - all_variables[k] = self.__cleanup(ev, 0) if self.cup_cleanup_variables else ev + all_variables[k] = self.__cleanup(ev, 0) if self.state.options.cup_cleanup_variables else ev if self.debug_level == DEBUG_LEVEL.full: - self.logger.debug(self.format_output(f"All variables: {all_variables}")) + self.logger.debug(format_output(f"All variables: {all_variables}")) # Insertions in the negative prompt if self.debug_level == DEBUG_LEVEL.full: - self.logger.debug(self.format_output(f"New negative additions: {p_processor.add_at}")) - self.logger.debug(self.format_output(f"New negative indexes: {n_processor.insertion_at}")) + self.logger.debug(format_output(f"New negative additions: {p_processor.add_at}")) + self.logger.debug(format_output(f"New negative indexes: {n_processor.insertion_at}")) negative_prompt = self.__add_to_insertion_points( negative_prompt, p_processor.add_at["insertion_point"], n_processor.insertion_at ) @@ -1107,19 +948,19 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in foundP = bool(p_processor.detectedWildcards) foundNP = bool(n_processor.detectedWildcards) if foundP or foundNP: - if self.wil_ifwildcards == self.IFWILDCARDS_CHOICES.stop: + if self.state.options.wil_ifwildcards == IFWILDCARDS_CHOICES.stop: self.logger.error("Found unprocessed wildcards!") else: self.logger.info("Found unprocessed wildcards.") ppwl = ", ".join(p_processor.detectedWildcards) npwl = ", ".join(n_processor.detectedWildcards) if foundP: - self.logger.error(self.format_output(f"In the positive prompt: {ppwl}")) + self.logger.error(format_output(f"In the positive prompt: {ppwl}")) if foundNP: - self.logger.error(self.format_output(f"In the negative prompt: {npwl}")) - if self.wil_ifwildcards == self.IFWILDCARDS_CHOICES.warn: + self.logger.error(format_output(f"In the negative prompt: {npwl}")) + if self.state.options.wil_ifwildcards == IFWILDCARDS_CHOICES.warn: prompt = self.WILDCARD_WARNING + prompt - elif self.wil_ifwildcards == self.IFWILDCARDS_CHOICES.stop: + elif self.state.options.wil_ifwildcards == IFWILDCARDS_CHOICES.stop: raise PPPInterrupt( "Found unprocessed wildcards!", self.WILDCARD_STOP.format(ppwl) if foundP else "", @@ -1156,25 +997,25 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in try: if seed == -1: seed = np.random.randint(0, 2**32, dtype=np.int64) - self.rng = np.random.default_rng(seed & 0xFFFFFFFF) prompt = original_prompt negative_prompt = original_negative_prompt - self.debug_level = DEBUG_LEVEL(self.options.get("debug_level", DEBUG_LEVEL.none.value)) if self.debug_level != DEBUG_LEVEL.none: - self.logger.info(f"System variables: {self.system_variables}") + self.logger.info(f"System variables: {self.state.system_variables}") self.logger.info(f"Input seed: {seed}") - self.logger.info(self.format_output(f"Input prompt: {prompt}")) - self.logger.info(self.format_output(f"Input negative_prompt: {negative_prompt}")) + self.logger.info(format_output(f"Input prompt: {prompt}")) + self.logger.info(format_output(f"Input negative_prompt: {negative_prompt}")) t1 = time.monotonic_ns() - prompt, negative_prompt, all_variables = self.__processprompts(prompt, negative_prompt) + prompt, negative_prompt, all_variables = self.__processprompts( + np.random.default_rng(seed & 0xFFFFFFFF), prompt, negative_prompt + ) t2 = time.monotonic_ns() if self.debug_level != DEBUG_LEVEL.none: - self.logger.info(self.format_output(f"Result prompt: {prompt}")) - self.logger.info(self.format_output(f"Result negative_prompt: {negative_prompt}")) + self.logger.info(format_output(f"Result prompt: {prompt}")) + self.logger.info(format_output(f"Result negative_prompt: {negative_prompt}")) self.logger.info(f"Process prompt pair time: {(t2 - t1) / 1_000_000_000:.3f} seconds") # if self.debug_level != DEBUG_LEVEL.none: - # self.logger.debug(f"Wildcards memory usage: {self.wildcard_obj.__sizeof__()}") + # self.logger.debug(f"Wildcards memory usage: {self.state.wildcards_obj.__sizeof__()}") # Check for constructs not processed due to parsing problems fullcontent: str = prompt + negative_prompt if fullcontent.find("= 0: @@ -1196,1676 +1037,3 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in except Exception as e: # pylint: disable=broad-exception-caught self.logger.exception(e) return original_prompt, original_negative_prompt, all_variables - - def parse_prompt(self, prompt_description: str, prompt: str, parser: lark.Lark, raise_parsing_error: bool = False): - """ - Parses a prompt using the specified parser. - - Args: - prompt_description (str): The description of the prompt. - prompt (str): The prompt to be parsed. - parser (lark.Lark): The parser to be used. - raise_parsing_error (bool): Whether to raise a parsing error. - - Returns: - Tree: The parsed prompt. - """ - t1 = time.monotonic_ns() - parsed_prompt = None - try: - if self.debug_level == DEBUG_LEVEL.full: - self.logger.debug(self.format_output(f"Parsing {prompt_description}: '{escape_single_quotes(prompt)}'")) - parsed_prompt = parser.parse(prompt) - # we store the contents so we can use them later even if the meta position is not valid anymore - if isinstance(parsed_prompt, lark.Tree): - for n in parsed_prompt.iter_subtrees(): - if isinstance(n, lark.Tree): - if n.meta.empty: - n.meta.content = "" - else: - n.meta.content = prompt[n.meta.start_pos : n.meta.end_pos] - except lark.exceptions.UnexpectedInput: - if raise_parsing_error: - raise - self.logger.exception(self.format_output(f"Parsing failed on prompt!: {escape_single_quotes(prompt)}")) - t2 = time.monotonic_ns() - if self.debug_level == DEBUG_LEVEL.full: - self.logger.debug(f"Parse {prompt_description} time: {(t2 - t1) / 1_000_000_000:.3f} seconds") - if parsed_prompt: - self.logger.debug( - "Tree:\n" - + textwrap.indent( - re.sub( - r"\n$", - "", - parsed_prompt.pretty() if isinstance(parsed_prompt, lark.Tree) else parsed_prompt, - ), - " ", - ) - ) - return parsed_prompt - - def warn_or_stop(self, is_negative: bool, message: str, e: Exception = None): - if self.gen_onwarning == self.ONWARNING_CHOICES.stop: - raise PPPInterrupt( - message, - self.INVALID_CONTENT_STOP.format(message) if not is_negative else "", - self.INVALID_CONTENT_STOP.format(message) if is_negative else "", - ) from e - self.logger.warning(message) - - class TreeProcessor(lark.visitors.Interpreter): - """ - A class for interpreting and processing a tree generated by the prompt parser. - - Args: - ppp (PromptPostProcessor): The PromptPostProcessor object. - - Attributes: - add_at (dict): The dictionary to store the content to be added at different positions of the negative prompt. - insertion_at (list): The list of insertion points in the negative prompt. - detectedWildcards (list): The list of detected invalid wildcards or choices. - result (str): The final processed prompt. - """ - - def __init__(self, ppp: "PromptPostProcessor"): - super().__init__() - self.__ppp = ppp - self.AccumulatedShell = namedtuple("AccumulatedShell", ["type", "data"]) - self.NegTag = namedtuple("NegTag", ["start", "end", "content", "parameters", "shell"]) - self.__shell: list[self.AccumulatedShell] = [] # type: ignore - self.__negtags: list[self.NegTag] = [] # type: ignore - self.__already_processed: list[str] = [] - self.__is_negative = False - self.__wildcard_filters = {} - self.__seen_wildcards: list[str] = [] - self.add_at: dict = {"start": [], "insertion_point": [[] for x in range(10)], "end": []} - self.insertion_at: list[tuple[int, int]] = [None for x in range(10)] - self.detectedWildcards: list[str] = [] - self.result = "" - - def warn_or_stop(self, message: str, e: Exception = None): - self.__ppp.warn_or_stop(self.__is_negative, message, e) - - def start_visit(self, prompt_description: str, parsed_prompt: lark.Tree, is_negative: bool = False) -> str: - """ - Start the visit process. - - Args: - prompt_description (str): The description of the prompt. - parsed_prompt (Tree): The parsed prompt. - is_negative (bool): Whether the prompt is negative or not. - - Returns: - str: The processed prompt. - """ - t1 = time.monotonic_ns() - self.__is_negative = is_negative - if self.__ppp.debug_level != DEBUG_LEVEL.none: - self.__ppp.logger.info(f"Processing {prompt_description}...") - self.visit(parsed_prompt) - t2 = time.monotonic_ns() - if self.__ppp.debug_level != DEBUG_LEVEL.none: - self.__ppp.logger.info(f"Process {prompt_description} time: {(t2 - t1) / 1_000_000_000:.3f} seconds") - return self.result - - def __visit( - self, - node: lark.Tree | lark.Token | list[lark.Tree | lark.Token] | None, - restore_state: bool = False, - discard_content: bool = False, - ) -> str: - """ - Visit a node in the tree and process it or accumulate its value if it is a Token. - - Args: - node (Tree|Token|list): The node or list of nodes to visit. - restore_state (bool): Whether to restore the state after visiting the node. - discard_content (bool): Whether to discard the content of the node. - - Returns: - str: The result of the visit. - """ - backup_result = self.result - # if self.__ppp.debug_level == DEBUG_LEVEL.full: - # self.__ppp.logger.debug(f"Visiting node {node}.") - if restore_state: - # if self.__ppp.debug_level == DEBUG_LEVEL.full: - # self.__ppp.logger.debug("Backing up state before visiting.") - backup_shell = self.__shell.copy() - backup_negtags = self.__negtags.copy() - backup_already_processed = self.__already_processed.copy() - backup_add_at = self.add_at.copy() - backup_insertion_at = self.insertion_at.copy() - backup_detectedwildcards = self.detectedWildcards.copy() - backup_user_variables = self.__ppp.user_variables.copy() - backup_echoed_variables = self.__ppp.echoed_variables.copy() - if node is not None: - if isinstance(node, list): - for child in node: - self.__visit(child) - elif isinstance(node, lark.Tree): - self.visit(node) - elif isinstance(node, lark.Token): - self.result += node - len_backup = len(backup_result) - # if self.result[:len_backup] == backup_result: # this is only necessary if we call parse_prompt with a parser from "start", because it resets the result - added_result = self.result[len_backup:] - # else: - # added_result = self.result - if discard_content or restore_state: - self.result = backup_result - if restore_state: - # if self.__ppp.debug_level == DEBUG_LEVEL.full: - # self.__ppp.logger.debug("Restoring state after visiting.") - self.__shell = backup_shell - self.__negtags = backup_negtags - self.__already_processed = backup_already_processed - self.add_at = backup_add_at - self.insertion_at = backup_insertion_at - self.detectedWildcards = backup_detectedwildcards - self.__ppp.user_variables = backup_user_variables - self.__ppp.echoed_variables = backup_echoed_variables - return added_result - - def __get_original_node_content(self, node: lark.Tree | lark.Token, default=None) -> str: - """ - Get the original content of a node. - - Args: - node (Tree|Token): The node to get the content from. - default: The default value to return if the content is not found. - - Returns: - str: The original content of the node. - """ - return ( - node.meta.content - if hasattr(node, "meta") and node.meta is not None and not node.meta.empty - else default - ) - - def __get_user_variable_value(self, name: str, evaluate=True, visit=False) -> str: - """ - Get the value of a user variable. - - Args: - name (str): The name of the user variable. - evaluate (bool): Whether to evaluate the variable. - visit (bool): Whether to also visit the variable (add to result). - - Returns: - str: The value of the user variable. - """ - v = self.__ppp.user_variables.get(name, None) - if v is not None: - visited = False - if isinstance(v, lark.Tree): - if evaluate: - v = self.__visit(v, not visit) - visited = visit - else: - v = self.__get_original_node_content(v, "not evaluated yet") - if visit and not visited: - self.result += v - return v - - def get_final_user_variable(self, name: str) -> str: - return self.__get_user_variable_value(name, True, False) - - def __set_user_variable_value(self, name: str, value: str): - """ - Set the value of a user variable. - - Args: - name (str): The name of the user variable. - value (str): The value to be set. - """ - self.__ppp.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.__ppp.user_variables: - del self.__ppp.user_variables[name] - - def __debug_end(self, construct: str, start_result: str, duration: int, info=None): - """ - Log the end of a construct processing. - - Args: - construct (str): The name of the construct. - start_result (str): The initial result. - duration (int): The duration of the processing in ns. - info: Additional information to log. - """ - if self.__ppp.debug_level == DEBUG_LEVEL.full: - info = f"({info}) " if info is not None and info != "" else "" - output = self.result[len(start_result) :] - if output != "": - output = f" >> '{escape_single_quotes(output)}'" - self.__ppp.logger.debug( - self.__ppp.format_output( - f"TreeProcessor.{construct} {info}({duration / 1_000_000_000:.3f} seconds){output}" - ) - ) - - def __resolve_cond_value(self, c: str): - """Resolve a condition value: try int first, fall back to variable lookup.""" - try: - return int(c) - except ValueError: - # Bare identifier - resolve as variable reference - if c.startswith("_"): - val = self.__ppp.system_variables.get(c, None) - if val is None: - val = "" - self.warn_or_stop(f"Unknown system variable {c}") - else: - val = self.__get_user_variable_value(c) - if val is None: - val = "" - self.warn_or_stop(f"Unknown user variable {c}") - return val.lower() if isinstance(val, str) else val - - def __eval_basiccondition(self, cond_var: str, cond_comp: str, cond_value: str | list[str]) -> bool: - """ - Evaluate a condition based on the given variable, comparison, and value. - - Args: - cond_var (str): The variable to be compared. - cond_comp (str): The comparison operator. - cond_value (str or list[str]): The value to be compared with. - - Returns: - bool: The result of the condition evaluation. - """ - if cond_var.lower() == "false": - var_value = "false" - elif cond_var.lower() == "true": - var_value = "true" - elif cond_var.startswith("_"): # system variable - var_value = self.__ppp.system_variables.get(cond_var, None) - if var_value is None: - var_value = "" - self.warn_or_stop(f"Unknown system variable {cond_var}") - else: # user variable - var_value = self.__get_user_variable_value(cond_var) - if var_value is None: - var_value = "" - self.warn_or_stop(f"Unknown user variable {cond_var}") - if isinstance(var_value, str): - var_value = var_value.lower() - if isinstance(cond_value, list): - comp_ops = { - "contains": lambda x, y: y in x, - "in": lambda x, y: x == y, - } - else: - cond_value = [cond_value] - comp_ops = { - "eq": lambda x, y: x == y, - "ne": lambda x, y: x != y, - "gt": lambda x, y: x > y, - "lt": lambda x, y: x < y, - "ge": lambda x, y: x >= y, - "le": lambda x, y: x <= y, - "contains": lambda x, y: y in x, - "truthy": lambda x, y: bool(x), - } - if cond_comp not in comp_ops: - return False - cond_value_adjusted = list( - ( - c[1:-1].lower() - if c.startswith('"') or c.startswith("'") - else ( - True - if c.lower() == "true" - else False if c.lower() == "false" or c == "" else self.__resolve_cond_value(c) - ) - ) - for c in cond_value - ) - result = False - for c in cond_value_adjusted: - if isinstance(c, str): - var_value_adjusted = var_value - elif isinstance(c, bool) and var_value != "false" and var_value != "" and var_value is not False: - var_value_adjusted = True - elif isinstance(c, bool) and (var_value != "true" or var_value is False): - var_value_adjusted = False - else: - try: - var_value_adjusted = int(var_value) - except (ValueError, TypeError): - self.warn_or_stop( - f"Cannot convert variable value '{escape_single_quotes(var_value)}' to integer for comparison" - ) - return False - result = comp_ops[cond_comp](var_value_adjusted, c) - if result: - break - return result - - def __eval_condition(self, condition: lark.Tree) -> bool: - """ - Evaluate an if condition based on the given condition tree. - - Args: - condition (Node): The condition tree to be evaluated. - - Returns: - bool: The result of the if condition evaluation. - """ - # self.__ppp.logger.debug(f"__eval_condition {condition.data}") - if condition.data == "operation_and": - cond_result = True - for c in condition.children: - cond_result = cond_result and self.__eval_condition(c) - if not cond_result: - break - elif condition.data == "operation_or": - cond_result = False - for c in condition.children: - cond_result = cond_result or self.__eval_condition(c) - if cond_result: - break - elif condition.data == "operation_not": - cond_result = not self.__eval_condition(condition.children[0]) - else: # truthy_operand / comparison_simple_value / comparison_list_value - # we get the name of the variable - cond_var = str(condition.children[0]) - poscomp = 1 - invert = False - if poscomp >= len(condition.children): - # no condition, just a variable - cond_comp = "truthy" - cond_value = "true" - else: - # we get the comparison (with possible not) and the value - cond_comp = str(condition.children[poscomp]) - if cond_comp == "not": - invert = not invert - poscomp += 1 - cond_comp = str(condition.children[poscomp]) - poscomp += 1 - cond_value_node = condition.children[poscomp] - cond_value = ( - list(str(v) for v in cond_value_node.children) - if isinstance(cond_value_node, (lark.Tree, list)) - else str(cond_value_node) if isinstance(cond_value_node, lark.Token) else cond_value_node - ) - cond_result = self.__eval_basiccondition(cond_var, cond_comp, cond_value) - if invert: - cond_result = not cond_result - return cond_result - - def promptcomp(self, tree: lark.Tree): - """ - Process a prompt composition construct in the tree. - """ - start_result = self.result - t1 = time.monotonic_ns() - self.__visit(tree.children[0]) - and_processing = self.__ppp.host_config.get("and", "ok") - if len(tree.children) > 1: - and_replacements = { - "eol": ("replaced with EOL", "\n"), - "comma": ("replaced with COMMA", ", "), - "remove": ("removed", " "), - } - if tree.children[1] is not None: - self.result += f":{tree.children[1]}" - for i in range(2, len(tree.children), 3): - if and_processing in and_replacements.keys(): - self.result = ( - self.result.rstrip() - + and_replacements[and_processing][1] - + self.__visit(tree.children[i + 1], False, True).lstrip() - ) - if self.__ppp.debug_level == DEBUG_LEVEL.full: - self.__ppp.logger.debug(f"AND construct {and_replacements[and_processing][0]}") - elif and_processing == "error": - self.warn_or_stop("AND constructs are not allowed!") - else: # and_processing == "ok": - if self.__ppp.cup_ands: - self.result = re.sub(r"[, ]+$", "\n" if self.__ppp.cup_ands_eol else " ", self.result) - if self.result[-1:].isalnum(): # add space if needed - self.result += " " - self.result += "AND" - added_result = self.__visit(tree.children[i + 1], False, True) - if self.__ppp.cup_ands: - added_result = re.sub(r"^[, ]+", " ", added_result) - if added_result[0:1].isalnum(): # add space if needed - added_result = " " + added_result - self.result += added_result - if tree.children[i + 2] is not None: - self.result += f":{tree.children[i+2]}" - t2 = time.monotonic_ns() - self.__debug_end("promptcomp", start_result, t2 - t1) - - def scheduled(self, tree: lark.Tree): - """ - Process a scheduling construct in the tree and add it to the accumulated shell. - """ - start_result = self.result - t1 = time.monotonic_ns() - before = tree.children[0] - after = tree.children[-2] - pos_str = tree.children[-1] - pos = float(pos_str) - if pos >= 1: - pos = int(pos) - scheduling_processing = self.__ppp.host_config.get("scheduling", "ok") - if scheduling_processing == "before": - if self.__ppp.debug_level == DEBUG_LEVEL.full: - self.__ppp.logger.debug("Scheduling construct removed, taking before option") - if before is not None: - self.__visit(before) - elif scheduling_processing == "after": - if self.__ppp.debug_level == DEBUG_LEVEL.full: - self.__ppp.logger.debug("Scheduling construct removed, taking after option") - if after is not None: - self.__visit(after) - elif scheduling_processing == "first": - if self.__ppp.debug_level == DEBUG_LEVEL.full: - self.__ppp.logger.debug("Scheduling construct removed, taking first option") - if before is not None: - self.__visit(before) - elif after is not None: - self.__visit(after) - elif scheduling_processing == "remove": - if self.__ppp.debug_level == DEBUG_LEVEL.full: - self.__ppp.logger.debug("Scheduling construct removed") - elif scheduling_processing == "error": - self.warn_or_stop("Scheduling constructs are not allowed!") - else: # scheduling_processing == "ok" - # self.__shell.append(self.AccumulatedShell("sc", pos)) - self.result += "[" - if before is not None: - if self.__ppp.debug_level == DEBUG_LEVEL.full: - self.__ppp.logger.debug(f"Shell scheduled before with position {pos}") - self.__shell.append(self.AccumulatedShell("scb", pos)) - self.__visit(before) - self.__shell.pop() - if self.__ppp.debug_level == DEBUG_LEVEL.full: - self.__ppp.logger.debug(f"Shell scheduled after with position {pos}") - self.__shell.append(self.AccumulatedShell("sca", pos)) - self.result += ":" - self.__visit(after) - self.__shell.pop() - if self.__ppp.cup_emptyconstructs and re.fullmatch(re.escape(start_result) + r"\[:\s*", self.result): - self.result = start_result - else: - self.result += f":{pos_str}]" - # self.__shell.pop() - t2 = time.monotonic_ns() - self.__debug_end("scheduled", start_result, t2 - t1, pos_str) - - def alternate(self, tree: lark.Tree): - """ - Process an alternation construct in the tree and add it to the accumulated shell. - """ - start_result = self.result - t1 = time.monotonic_ns() - alternation_processing = self.__ppp.host_config.get("alternation", "ok") - if alternation_processing == "first": - if self.__ppp.debug_level == DEBUG_LEVEL.full: - self.__ppp.logger.debug("Alternation construct removed, taking first option") - self.__visit(tree.children[0]) - elif alternation_processing == "remove": - if self.__ppp.debug_level == DEBUG_LEVEL.full: - self.__ppp.logger.debug("Alternation construct removed") - elif alternation_processing == "error": - self.warn_or_stop("Alternation constructs are not allowed!") - else: # alternation_processing == "ok" - # self.__shell.append(self.AccumulatedShell("al", len(tree.children))) - self.result += "[" - for i, opt in enumerate(tree.children): - if self.__ppp.debug_level == DEBUG_LEVEL.full: - self.__ppp.logger.debug(f"Shell alternate option {i+1}") - self.__shell.append(self.AccumulatedShell("alo", {"pos": i + 1, "len": len(tree.children)})) - if i > 0: - self.result += "|" - self.__visit(opt) - self.__shell.pop() - self.result += "]" - if self.__ppp.cup_emptyconstructs and re.fullmatch(re.escape(start_result) + r"\[\s*\]", self.result): - self.result = start_result - # self.__shell.pop() - t2 = time.monotonic_ns() - self.__debug_end("alternate", start_result, t2 - t1) - - def attention(self, tree: lark.Tree): - """ - Process a attention change construct in the tree and add it to the accumulated shell. - """ - start_result = self.result - t1 = time.monotonic_ns() - # weight_kind: -1: remove, 0=none, 1=decrease, 2=increase, 3=specific - if len(tree.children) == 2: - weight_str = tree.children[-1] - if weight_str is not None: - weight_kind = 3 # specific weight - weight = float(weight_str) - else: - weight_kind = 2 # increase attention - weight = 1.1 - weight_str = "1.1" - else: - weight_kind = 1 # decrease attention - weight = 0.9 - weight_str = "0.9" - if self.__ppp.debug_level == DEBUG_LEVEL.full: - self.__ppp.logger.debug(f"Shell attention with weight {weight}") - current_tree = tree.children[0] - if self.__ppp.cup_mergeattention: - while isinstance(current_tree, lark.Tree) and current_tree.data == "attention": - # we merge the weights - if len(current_tree.children) == 2: - inner_weight = current_tree.children[-1] - if inner_weight is not None: - inner_weight = float(inner_weight) - else: - inner_weight = 1.1 - else: - inner_weight = 0.9 - weight *= inner_weight - current_tree = current_tree.children[0] - weight = math.floor(weight * 100) / 100 # we round to 2 decimals - weight_str = f"{weight:.2f}".rstrip("0").rstrip(".") - if weight_str == "0.9": - weight_kind = 1 - elif weight_str == "1.1": - weight_kind = 2 - else: - weight_kind = 3 - attention_processing = self.__ppp.host_config.get("attention", "ok") - if attention_processing == "parentheses": - if weight_kind == 1: - weight_kind = 3 - weight_str = "0.9" - if self.__ppp.debug_level == DEBUG_LEVEL.full: - self.__ppp.logger.debug("Converted to parentheses format") - elif attention_processing == "disable": - weight_kind = 0 - if self.__ppp.debug_level == DEBUG_LEVEL.full: - self.__ppp.logger.debug("Attention construct disabled") - elif attention_processing == "remove": - weight_kind = -1 - if self.__ppp.debug_level == DEBUG_LEVEL.full: - self.__ppp.logger.debug("Attention construct removed") - elif attention_processing == "error": - self.warn_or_stop("Attention constructs are not allowed!") - # else: attention_processing == "ok": - if weight_kind == -1: - # we just ignore the attention construct - pass - elif weight_kind == 0: - # we just visit the content without adding any attention - self.__visit(current_tree) - else: - self.__shell.append(self.AccumulatedShell("at", (weight_kind, weight_str))) - if weight_kind == 1: - starttag = "[" - self.result += starttag - self.__visit(current_tree) - endtag = "]" - elif weight_kind == 2: - starttag = "(" - self.result += starttag - self.__visit(current_tree) - endtag = ")" - else: # weight_kind == 3 - starttag = "(" - self.result += starttag - self.__visit(current_tree) - endtag = f":{weight_str})" - if self.__ppp.cup_emptyconstructs and re.fullmatch( - re.escape(start_result + starttag) + r"\s*", self.result - ): - self.result = start_result - else: - self.result += endtag - self.__shell.pop() - t2 = time.monotonic_ns() - self.__debug_end("attention", start_result, t2 - t1, weight_str) - - def commandstn(self, tree: lark.Tree): - """ - Process a send to negative command in the tree and add it to the list of negative tags. - """ - start_result = self.result - info = None - t1 = time.monotonic_ns() - if not self.__is_negative: - negtagparameters = tree.children[0] - if negtagparameters is not None: - parameters = str(negtagparameters) - else: - parameters = "" - content = self.__visit(tree.children[1::], False, True) - self.__negtags.append( - self.NegTag(len(self.result), len(self.result), content, parameters, self.__shell.copy()) - ) - info = f"with {escape_single_quotes(parameters) or 'no parameters'} : {escape_single_quotes(content)}" - else: - self.warn_or_stop("Ignored negative command in negative prompt") - self.__visit(tree.children[1::]) - t2 = time.monotonic_ns() - self.__debug_end("commandstn", start_result, t2 - t1, info) - - def commandstni(self, tree: lark.Tree): - """ - Process a send to negative insertion point command in the tree and add it to the list of negative tags. - """ - start_result = self.result - info = None - t1 = time.monotonic_ns() - if self.__is_negative: - negtagparameters = tree.children[0] - if negtagparameters is not None: - parameters = str(negtagparameters) - else: - parameters = "" - self.__negtags.append( - self.NegTag(len(self.result), len(self.result), "", parameters, self.__shell.copy()) - ) - info = f"with {parameters or 'no parameters'}" - else: - self.warn_or_stop("Ignored negative insertion point command in positive prompt") - t2 = time.monotonic_ns() - self.__debug_end("commandstni", start_result, t2 - t1, info) - - def __varset( - self, - command: str, - variable: str, - modifiers: lark.Tree | None, - content: lark.Tree | None, - ): - """ - Process a generic set command in the tree. - """ - t1 = time.monotonic_ns() - start_result = self.result - if variable.startswith("_"): - self.warn_or_stop( - f"Invalid variable name '{escape_single_quotes(variable)}' detected! System variables cannot be set." - ) - return - info = variable - value_description = self.__get_original_node_content(content, None) - value = content - modifiers_str: list[str] = [str(m) for m in modifiers.children] if modifiers is not None else [] - if any(item in modifiers_str for item in ["+", "add"]): - info += f" += '{escape_single_quotes(value_description or '')}'" - raw_oldvalue = self.__ppp.user_variables.get(variable, None) - if raw_oldvalue is None: - newvalue = value - self.warn_or_stop(f"Unknown variable {escape_single_quotes(variable)}") - elif isinstance(raw_oldvalue, str): - newvalue = lark.Tree( - lark.Token("RULE", "varvalue"), - [lark.Token("plain", raw_oldvalue), value], - # Meta should be {"content": raw_oldvalue + value}, - ) - else: - newvalue = lark.Tree( - lark.Token("RULE", "varvalue"), - [raw_oldvalue, value], - # Meta should be {"content": raw_oldvalue.meta.content + value.meta.content}, - ) - elif any(item in modifiers_str for item in ["?", "ifundefined"]): - info += f" ?= '{escape_single_quotes(value_description or '')}'" - raw_oldvalue = self.__ppp.user_variables.get(variable, None) - if raw_oldvalue is None: - newvalue = value - else: - info += " (not set)" - newvalue = None - else: - newvalue = value - if newvalue is not None: - if any(item in modifiers_str for item in ["!", "evaluate"]): - newvalue = self.__visit(newvalue, False, True) - info += " =! " - else: - info += " = " - self.__set_user_variable_value(variable, newvalue) - currentvalue = self.__get_user_variable_value(variable, False) - if currentvalue is None: - info += "not evaluated yet" - else: - info += f"'{escape_single_quotes(currentvalue)}'" - t2 = time.monotonic_ns() - self.__debug_end(command, start_result, t2 - t1, info) - - def variableset(self, tree: lark.Tree): - """ - Process a DP set variable command in the tree and add it to the dictionary of variables. - """ - modifiers = tree.children[1] or lark.Tree(lark.Token("RULE", "variablesetmodifiers"), []) - immediate = tree.children[2] - if immediate is not None: - modifiers.children = modifiers.children.copy() - modifiers.children.append(immediate) - self.__varset("variableset", str(tree.children[0]), modifiers, tree.children[3]) - - def commandset(self, tree: lark.Tree): - """ - Process a set command in the tree and add it to the dictionary of variables. - """ - self.__varset("commandset", str(tree.children[0]), tree.children[1], tree.children[2]) - - def __varecho(self, command: str, variable: str, default: lark.Tree | None): - """ - Process a generic echo command in the tree. - """ - t1 = time.monotonic_ns() - start_result = self.result - default_value = None - # if default is not None: - # default_value = self.__visit(default, True) # for log - value = self.__get_user_variable_value(variable, True, True) - if value is None: - if default is not None: - if self.__ppp.debug_level == DEBUG_LEVEL.full: - self.__ppp.logger.debug( - f"Variable '{escape_single_quotes(variable)}' not found, using default value" - ) - v = self.__visit(default, False, True) - self.result += v - default_value = v - self.__ppp.echoed_variables[variable] = v - else: - self.warn_or_stop(f"Unknown variable {escape_single_quotes(variable)}") - default_value = "" - self.__ppp.echoed_variables[variable] = "" - else: - self.__ppp.echoed_variables[variable] = value - t2 = time.monotonic_ns() - info = variable - if default_value is not None: - info += f" with default '{escape_single_quotes(default_value)}'" - self.__debug_end(command, start_result, t2 - t1, info) - - def variableuse(self, tree: lark.Tree): - """ - Process a DP use variable command in the tree. - """ - self.__varecho("variableuse", str(tree.children[0]), tree.children[1] if len(tree.children) > 1 else None) - - def commandecho(self, tree: lark.Tree): - """ - Process an echo command in the tree. - """ - self.__varecho("commandecho", str(tree.children[0]), tree.children[1] if len(tree.children) > 1 else None) - - def commandif(self, tree: lark.Tree): - """ - Process an if command in the tree. - """ - t1 = time.monotonic_ns() - start_result = self.result - for i, n in enumerate(tree.children): - content = n.children[-1] - if len(n.children) == 2: # its not an else - # has a condition - condition = n.children[0] - c = self.__get_original_node_content(condition, f"condition {i}") - if self.__eval_condition(condition): - self.__visit(content) - t2 = time.monotonic_ns() - self.__debug_end("commandif", start_result, t2 - t1, c) - return - else: # its an else - self.__visit(content) - t2 = time.monotonic_ns() - self.__debug_end("commandif", start_result, t2 - t1, "else") - return - - def commandext(self, tree: lark.Tree): - """ - Process an extranetwork command in the tree. - """ - t1 = time.monotonic_ns() - start_result = self.result - extnet = "(ignored)" - if not self.__ppp.rem_removeextranetworktags: - extnet_type: str = (tree.children[0].children[0] or "") + str(tree.children[0].children[1]) - is_mapping = extnet_type.startswith("$") - if is_mapping: - extnet_type = extnet_type[1:] - extnet_id: str = str(tree.children[1]) - if extnet_id.startswith("'") or extnet_id.startswith('"'): - extnet_id = extnet_id[1:-1] - extnet_id = re.sub(r"\\(.)", r"\1", extnet_id) # so we can escape some special characters - parameters: str = "" - parameters_defaulted = False - if tree.children[2]: - parameters = str(tree.children[2]) - elif extnet_type in ("lora", "hypernet"): - parameters = "1" - parameters_defaulted = True - if parameters.startswith("'") or parameters.startswith('"'): - parameters = parameters[1:-1] - parameters_is_number = bool(re.match(r"^[-+]?\d*\.?\d+$", parameters or "")) - condition = tree.children[3] - if not condition or self.__eval_condition(condition): - extnet_id = f"{extnet_type}:{extnet_id}" - triggers = tree.children[4] if len(tree.children) > 4 else None - extra_triggers = None - compiled_extra_triggers = None - if is_mapping: - found = self.__ppp.extranetwork_mappings_obj.cached_mappings.get(extnet_id, None) - # we assume the conditions do not change inside the prompt - found_in_cache = found is not None - if found is None: - found_mappings: list[PPPENMappingVariant] = [] - else_mapping = None - if self.__ppp.extranetwork_mappings_obj: - enmapping = self.__ppp.extranetwork_mappings_obj.extranetwork_mappings.get( - extnet_id, None - ) - if enmapping: - for v in enmapping.variants: - if v.condition: - try: - cnd = self.__ppp.parse_prompt( - "condition", v.condition, self.__ppp.parser_condition, True - ) - except lark.exceptions.UnexpectedInput as e: - self.warn_or_stop( - f"Error parsing condition '{escape_single_quotes(v.condition)}' in extranetwork mapping '{escape_single_quotes(extnet_id)}'! : {e.__class__.__name__}", - e, - ) - cnd = None - else: - cnd = "True" - if cnd is not None and (cnd == "True" or self.__eval_condition(cnd)): - if v.condition: - found_mappings.append(v) - else: - else_mapping = v - if found_mappings: - found = found_mappings[ - self.__ppp.rng.choice( - len(found_mappings), - p=[v.weight or 1 for v in found_mappings], - ) - ] - else: - found = else_mapping - self.__ppp.extranetwork_mappings_obj.cached_mappings[extnet_id] = found - if found: - if found.name: - if not found_in_cache and self.__ppp.debug_level != DEBUG_LEVEL.none: - self.__ppp.logger.info( - f"Mapping extranetwork '{escape_single_quotes(extnet_id)}' to '{escape_single_quotes(extnet_type)}:{escape_single_quotes(found.name)}'" - ) - extnet_id = f"{extnet_type}:{found.name}" - f_parameters = found.parameters - if not f_parameters and extnet_type in ("lora", "hypernet"): - f_parameters = "1" - found_parameters_is_number = True - else: - found_parameters_is_number = f_parameters and bool( - re.match(r"^[-+]?\d*\.?\d+$", str(f_parameters) or "") - ) - if found_parameters_is_number and parameters_is_number: - parameters = f"{float(f_parameters) * float(parameters):.2f}".rstrip("0").rstrip( - "." - ) - elif f_parameters is not None and parameters_defaulted: - parameters = f_parameters - elif found.triggers: - if not found_in_cache and self.__ppp.debug_level != DEBUG_LEVEL.none: - self.__ppp.logger.info( - f"Mapping extranetwork '{escape_single_quotes(extnet_id)}' to just triggers" - ) - extnet_id = None - else: - if not found_in_cache and self.__ppp.debug_level != DEBUG_LEVEL.none: - self.__ppp.logger.info( - f"Mapping extranetwork '{escape_single_quotes(extnet_id)}' to nothing" - ) - extnet_id = None - if found.triggers: - extra_triggers = ", ".join(found.triggers) - try: - compiled_extra_triggers = self.__ppp.parse_prompt( - "triggers", extra_triggers, self.__ppp.parser_content, True - ) - except lark.exceptions.UnexpectedInput as e: - self.warn_or_stop( - f"Error parsing triggers '{escape_single_quotes(extra_triggers)}' in extranetwork mapping '{escape_single_quotes(extnet_id)}'! : {e.__class__.__name__}", - e, - ) - compiled_extra_triggers = None - else: - self.warn_or_stop(f"Extranetwork mapping '{escape_single_quotes(extnet_id)}' not found!") - if extnet_id: - extnet = f"<{extnet_id}:{parameters}>" - self.result += extnet - elif triggers or compiled_extra_triggers: - extnet = "(only triggers)" - if triggers or compiled_extra_triggers: - if extnet_id: - if not self.__ppp.cup_extranetworktags: - self.result += " " - else: - self.result += ", " - if triggers: - self.result += self.__visit(triggers, True, True) - if compiled_extra_triggers: - if triggers: - self.result += ", " - self.result += self.__visit(compiled_extra_triggers, True, True) - if triggers or compiled_extra_triggers: - self.result += ", " - t2 = time.monotonic_ns() - self.__debug_end("commandext", start_result, t2 - t1, extnet) - - def commandsetwcdeffilter(self, tree: lark.Tree): - """ - Process a setwcdeffilter (Set Wildcard Default Filter) command in the tree. - """ - t1 = time.monotonic_ns() - start_result = self.result - wildcard_key: str = self.__visit(tree.children[0].children[1], False, True) - selected_wildcards = [x.key for x in self.__ppp.wildcard_obj.get_wildcards(wildcard_key)] - if not selected_wildcards: - self.warn_or_stop( - f"Wildcard '{escape_single_quotes(wildcard_key)}' not found for default filter setting!" - ) - else: - filter_object = tree.children[1].children[1] if tree.children[1] is not None else None - if filter_object is None: - for wc in selected_wildcards: - if self.__ppp.debug_level == DEBUG_LEVEL.full: - self.__ppp.logger.debug(f"Removed default filter for wildcard '{escape_single_quotes(wc)}'") - self.__ppp.wildcard_obj.set_wildcard_default_filter(wc, None) - else: - filter_specifier = self.__extract_filter_specifiers(filter_object) - for wc in selected_wildcards: - if self.__ppp.debug_level == DEBUG_LEVEL.full: - self.__ppp.logger.debug(f"Set default filter for wildcard '{escape_single_quotes(wc)}'") - self.__ppp.wildcard_obj.set_wildcard_default_filter(wc, filter_specifier) - t2 = time.monotonic_ns() - self.__debug_end("commandsetwcdeffilter", start_result, t2 - t1) - - def extranetworktag(self, tree: lark.Tree): - """ - Process an extra network construct in the tree. - """ - t1 = time.monotonic_ns() - start_result = self.result - if not self.__ppp.rem_removeextranetworktags: - self.result += f"<{tree.children[0]}" - self.__visit(tree.children[1]) - self.result += ">" - t2 = time.monotonic_ns() - self.__debug_end("extranetworktag", start_result, t2 - t1) - - def __get_choices_internal_get( - self, - choice_values: list[dict], - filter_specifier: Optional[list[list[str]]] = None, - wildcard_key: str = None, - ) -> list[dict]: - msg_where = f"wildcard '{escape_single_quotes(wildcard_key)}'" if wildcard_key else "choices" - if filter_specifier is not None: - filtered_choice_values = [] - for i, c in enumerate(choice_values): - passes = False - for or_ in filter_specifier: - tmp_pass = True - for and_ in or_: - if and_.isdecimal(): - if int(and_) != i: - tmp_pass = False - break - elif re.match(r"^\d+-\d+$", and_): - try: - start, end = map(int, and_.split("-", 1)) - if not (start <= i <= end): - tmp_pass = False - break - except ValueError: - tmp_pass = False - break - elif and_.lower() not in c.get("labels", []): - tmp_pass = False - break - if tmp_pass: - passes = True - break - if passes: - filtered_choice_values.append(c) - if not filtered_choice_values: - self.warn_or_stop( - f"Wildcard filter specifier '{escape_single_quotes(','.join(['+'.join(y for y in x) for x in filter_specifier]))}' found no matches in choices for wildcard '{escape_single_quotes(wildcard_key)}'!" - ) - else: - filtered_choice_values = choice_values.copy() - expanded_choice_values = [] - for i, c in enumerate(filtered_choice_values): - if c.get("command", False): - content_text = self.__visit(c.get("content", ""), False, True).strip() - (cmd, cmd_args) = content_text.split() - if cmd == "include": - wcs = self.__ppp.wildcard_obj.get_wildcards(cmd_args) - if not wcs: - self.warn_or_stop( - f"Not found included wildcard '{escape_single_quotes(cmd_args)}' at {msg_where}!" - ) - c_weight = float(c.get("weight", 1.0)) - for wc in wcs: - if wc.key in self.__seen_wildcards: - self.warn_or_stop( - f"Circular reference detected including wildcard '{escape_single_quotes(wc.key)}' at {msg_where} (chain starts at '{escape_single_quotes(self.__seen_wildcards[0])}')!" - ) - continue - self.__seen_wildcards.append(wc.key) - if self.__ppp.debug_level == DEBUG_LEVEL.full: - self.__ppp.logger.debug(f"Seen wildcard '{escape_single_quotes(wc.key)}'") - self.__ppp.logger.debug( - f"Including choices from wildcard '{escape_single_quotes(wc.key)}'" - ) - (_, choice_values) = self.__check_wildcard_initialization(wc) - if choice_values is not None: - ch_values = self.__get_choices_internal_get(choice_values, None, wc.key) - for cv in ch_values: - expanded_choice_values.append( - { - **cv, - "weight": float(cv.get("weight", 1.0) * c_weight), # we adjust the weight - } - ) - else: - self.warn_or_stop(f"Unsupported choice command '{escape_single_quotes(cmd)}' at {msg_where}!") - else: - expanded_choice_values.append(c) - return expanded_choice_values - - def __get_choices_internal_select( - self, - options: dict | None, - choice_values: list[dict], - filter_specifier: Optional[list[list[str]]] = None, - wildcard_key: str = None, - ) -> tuple[str, list[str], str, str]: - """ - Select choices based on the options. - - Args: - options (dict): The object representing the options construct. - choice_values (list[dict]): A list of choice objects. - filter_specifier (list[list[str]]): The filter specifier. - wildcard_key (str): The wildcard key if it is a wildcard. - - Returns: - tuple: A tuple containing the prefix, selected choices, separator and suffix - """ - seen_wildcards_len = len(self.__seen_wildcards) - if options is None: - options = {} - sampler: str = options.get("sampler", "~") - repeating: bool = options.get("repeating", False) - optional: bool = options.get("optional", False) - if "count" in options: - from_value = options["count"] - to_value = from_value - else: - from_value: int = options.get("from", 1) - to_value: int = options.get("to", 1) - separator: str = options.get("separator", self.__ppp.wil_choice_separator) - msg_where = f"wildcard '{escape_single_quotes(wildcard_key)}'" if wildcard_key else "choices" - if sampler != "~": - self.warn_or_stop(f"Unsupported sampler '{escape_single_quotes(sampler)}' at {msg_where} options!") - sampler = "~" - expanded_choice_values = self.__get_choices_internal_get(choice_values, filter_specifier, wildcard_key) - available_choices: list[dict] = [] - weights = [] - included_choices = 0 - excluded_choices = 0 - excluded_weights_sum = 0 - for i, c in enumerate(expanded_choice_values): - c["choice_index"] = i # we index them to later sort the results - weight = float(c.get("weight", 1.0)) - condition = c.get("if", None) - if weight > 0 and (condition is None or self.__eval_condition(condition)): - available_choices.append(c) - weights.append(weight) - included_choices += 1 - else: - weights.append(-1) - excluded_choices += 1 - excluded_weights_sum += weight - if excluded_choices > 0: # we need to redistribute the excluded weights - weights = [weight + excluded_weights_sum / included_choices for weight in weights if weight >= 0] - weights = np.array(weights) - weights /= weights.sum() # normalize weights - if available_choices: - if from_value < 0: - from_value = 1 - elif from_value > len(available_choices): - from_value = len(available_choices) - if to_value < 1: - to_value = 1 - elif (to_value > len(available_choices) and not repeating) or from_value > to_value: - to_value = len(available_choices) - num_choices = ( - self.__ppp.rng.integers(from_value, to_value, endpoint=True) - if from_value < to_value - else from_value - ) - else: - num_choices = 0 - if not optional and from_value > 0: - self.warn_or_stop(f"Not enough choices found for {msg_where}!") - if num_choices < 2: - repeating = False - if self.__ppp.debug_level == DEBUG_LEVEL.full: - self.__ppp.logger.debug( - self.__ppp.format_output( - f"Selecting {'optional ' if optional else ''}{'repeating ' if repeating else ''}{num_choices} choice" - + ("s" if num_choices != 1 else "") - + (f" and separating with '{escape_single_quotes(separator)}'" if num_choices > 1 else "") - ) - ) - if num_choices > 0: - selected_choices: list[dict] = ( - list(self.__ppp.rng.choice(available_choices, size=num_choices, p=weights, replace=repeating)) - if available_choices - else [] - ) - if self.__ppp.wil_keep_choices_order: - selected_choices = sorted(selected_choices, key=lambda x: x["choice_index"]) - selected_choices_text = [] - prefix: str = ( - self.__visit(options.get("prefix", None), False, True) - if options.get("prefix", None) is not None - else "" - ) - if prefix != "" and re.match(r"\w", prefix[-1]): - prefix += " " - for i, c in enumerate(selected_choices): - t1 = time.monotonic_ns() - choice_content_obj = c.get("content", c.get("text", None)) - if isinstance(choice_content_obj, str): - choice_content = choice_content_obj - else: - choice_content = self.__visit(choice_content_obj, False, True) - t2 = time.monotonic_ns() - if self.__ppp.debug_level == DEBUG_LEVEL.full: - self.__ppp.logger.debug( - f"Adding choice {i+1} ({(t2-t1) / 1_000_000_000:.3f} seconds):\n" - + textwrap.indent(re.sub(r"\n$", "", choice_content), " ") - ) - selected_choices_text.append(choice_content) - suffix: str = ( - self.__visit(options.get("suffix", None), False, True) - if options.get("suffix", None) is not None - else "" - ) - if suffix != "" and re.match(r"\w", suffix[0]): - suffix = " " + suffix - # remove comments - results = [re.sub(r"\s*#[^\n]*(?:\n|$)", "", r, flags=re.DOTALL) for r in selected_choices_text] - else: - prefix = "" - suffix = "" - results = [] - if self.__ppp.debug_level == DEBUG_LEVEL.full: - list_unseen = [f"'{escape_single_quotes(x)}'" for x in self.__seen_wildcards[seen_wildcards_len:]] - self.__ppp.logger.debug(f"Unseen wildcards: {', '.join(list_unseen)}") - self.__seen_wildcards = self.__seen_wildcards[:seen_wildcards_len] - return (prefix, results, separator, suffix) - - def __get_choices( - self, - options: dict | None, - choice_values: list[dict], - filter_specifier: Optional[list[list[str]]] = None, - wildcard_key: str = None, - ) -> str: - r = self.__get_choices_internal_select(options, choice_values, filter_specifier, wildcard_key) - if r[1]: - return r[0] + r[2].join(r[1]) + r[3] - return "" - - def __convert_choices_options(self, options: Optional[lark.Tree], is_wcdef: bool = False) -> dict: - """ - Convert the choices options to a dictionary. - - Args: - options (Tree): The choices options tree. - - Returns: - dict: The converted choices options. - """ - if options is None: - return None - options_dict = {} - if len(options.children) == 1: - if options.children[0] is not None: - options_dict["sampler"] = str(options.children[0]) - else: - if options.children[0] is not None: - options_dict["sampler"] = str(options.children[0].children[0]) - if options.children[1] is not None: - flags = str(options.children[1].children[0]) - options_dict["repeating"] = "r" in flags - options_dict["optional"] = "o" in flags - irange = 2 - idesc = 3 - isep = 4 - if options.children[irange] is not None: - if len(options.children[irange].children) == 1: - if options.children[irange].children[0] is not None: - options_dict["count"] = int(options.children[irange].children[0]) - else: - options_dict["from"] = ( - int(options.children[irange].children[0]) - if options.children[irange].children[0] is not None - else 1 - ) - options_dict["to"] = ( - int(options.children[irange].children[1]) - if options.children[irange].children[1] is not None - else 1 - ) - if is_wcdef: - if options.children[idesc] is not None: - options_dict["description"] = str(options.children[idesc].children[0])[1:-1] - else: - isep -= 1 # only wildcard definition options have a description - if options.children[isep] is not None: - options_dict["separator"] = self.__visit(options.children[isep], False, True) - if not options_dict: - options_dict = None - return options_dict - - def __convert_choice(self, choice: lark.Tree) -> dict: - """ - Convert the choice to a dictionary. - - Args: - choice (Tree): The choice tree. - - Returns: - dict: The converted choice. - """ - choice_dict = {} - choice_dict["command"] = choice.children[0] is not None - c_label_obj = choice.children[1] - choice_dict["labels"] = ( - [str(x).lower() for x in c_label_obj.children[1:-1]] if c_label_obj is not None else [] - ) - choice_dict["weight"] = float(choice.children[2].children[0]) if choice.children[2] is not None else 1.0 - choice_dict["if"] = choice.children[3].children[0] if choice.children[3] is not None else None - choice_dict["content"] = choice.children[-1] - return choice_dict - - def __check_wildcard_initialization(self, wildcard: PPPWildcard) -> tuple[dict | None, list[dict] | None]: - """ - Initializes a wildcard if it hasn't been yet. - - Args: - wildcard (PPPWildcard): The wildcard to check. - Returns: - tuple: A tuple containing the options and choice values of the wildcard. - """ - choice_values = wildcard.choices - options = wildcard.options - if choice_values is None: - t1 = time.monotonic_ns() - choice_values = [] - n = 0 - # we check the first choice to see if it is actually options - if isinstance(wildcard.unprocessed_choices[0], dict): - if self.__ppp.wildcard_obj.is_dict_wcdef_options(wildcard.unprocessed_choices[0]): - options = wildcard.unprocessed_choices[0] - prefix = options.get("prefix", None) - if prefix is not None and isinstance(prefix, str): - try: - options["prefix"] = self.__ppp.parse_prompt( - "choicevalue", prefix, self.__ppp.parser_choicevalue, True - ) - except lark.exceptions.UnexpectedInput as e: - self.warn_or_stop( - f"Error parsing choice prefix '{escape_single_quotes(prefix)}' in wildcard '{escape_single_quotes(wildcard.key)}'! : {e.__class__.__name__}", - e, - ) - suffix = options.get("suffix", None) - if suffix is not None and isinstance(suffix, str): - try: - options["suffix"] = self.__ppp.parse_prompt( - "choicevalue", suffix, self.__ppp.parser_choicevalue, True - ) - except lark.exceptions.UnexpectedInput as e: - self.warn_or_stop( - f"Error parsing choice suffix '{escape_single_quotes(suffix)}' in wildcard '{escape_single_quotes(wildcard.key)}'! : {e.__class__.__name__}", - e, - ) - n = 1 - else: - if wildcard.unprocessed_choices[0].endswith("$$"): - try: - options = self.__convert_choices_options( - self.__ppp.parse_prompt( - "as wildcard options", - wildcard.unprocessed_choices[0][:-2].strip(), - self.__ppp.parser_wcdefoptions, - True, - ), - True, - ) - n = 1 - except lark.exceptions.UnexpectedInput: - options = None - if options is None and self.__ppp.debug_level == DEBUG_LEVEL.full: - self.__ppp.logger.debug("Does not have options") - wildcard.options = options - # we process the choices - for cv in wildcard.unprocessed_choices[n:]: - if isinstance(cv, dict): - if self.__ppp.wildcard_obj.is_dict_choice_options(cv): - condition = cv.get("if", None) - if condition is not None and isinstance(condition, str): - try: - cv["if"] = self.__ppp.parse_prompt( - "condition", condition, self.__ppp.parser_condition, True - ) - except lark.exceptions.UnexpectedInput as e: - self.warn_or_stop( - f"Error parsing condition '{escape_single_quotes(condition)}' in wildcard '{escape_single_quotes(wildcard.key)}'! : {e.__class__.__name__}", - e, - ) - cv["if"] = None - content = cv.get("content", cv.get("text", None)) - cv["content"] = content - if "text" in cv: - del cv["text"] - if content is not None and isinstance(content, str): - try: - cv["content"] = self.__ppp.parse_prompt( - "choicevalue", content, self.__ppp.parser_choicevalue, True - ) - except lark.exceptions.UnexpectedInput as e: - self.warn_or_stop( - f"Error parsing choice content '{escape_single_quotes(content)}' in wildcard '{escape_single_quotes(wildcard.key)}'! : {e.__class__.__name__}", - e, - ) - cv["content"] = None - if cv["content"] is not None: - if self.__ppp.debug_level == DEBUG_LEVEL.full: - self.__ppp.logger.debug(f"Processed choice {cv}") - choice_values.append(cv) - else: - self.warn_or_stop( - f"Invalid choice {cv} in wildcard '{escape_single_quotes(wildcard.key)}'!" - ) - else: - self.warn_or_stop( - f"Invalid choice {cv} in wildcard '{escape_single_quotes(wildcard.key)}'!" - ) - else: - try: - choice_values.append( - self.__convert_choice( - self.__ppp.parse_prompt("choice", cv, self.__ppp.parser_choice, True) - ) - ) - except lark.exceptions.UnexpectedInput as e: - self.warn_or_stop( - f"Error parsing choice '{escape_single_quotes(cv)}' in wildcard '{escape_single_quotes(wildcard.key)}'! : {e.__class__.__name__}", - e, - ) - wildcard.choices = choice_values - t2 = time.monotonic_ns() - if self.__ppp.debug_level == DEBUG_LEVEL.full: - self.__ppp.logger.debug( - f"Processed choices for wildcard '{escape_single_quotes(wildcard.key)}' ({(t2-t1) / 1_000_000_000:.3f} seconds)" - ) - return (self.__clean_wildcard_options(options), choice_values) - - def __clean_wildcard_options(self, options: dict) -> dict: - if options is not None: - options.pop("description", None) # description is only for wildcard definitions, not for usage - if not options: - options = None - return options - - def wildcard(self, tree: lark.Tree): - """ - Process a wildcard construct in the tree. - """ - t1 = time.monotonic_ns() - seen_wildcards_len = len(self.__seen_wildcards) - start_result = self.result - applied_options = self.__clean_wildcard_options(self.__convert_choices_options(tree.children[0], False)) - wildcard_key: str = self.__visit(tree.children[1], False, True) - wc = self.__get_original_node_content(tree, f"?__{wildcard_key}__") - if self.__ppp.wil_process_wildcards: - if self.__ppp.debug_level == DEBUG_LEVEL.full: - self.__ppp.logger.debug(f"Processing wildcard: {wildcard_key}") - selected_wildcards = self.__ppp.wildcard_obj.get_wildcards(wildcard_key) - if not selected_wildcards: - self.detectedWildcards.append(wc) - self.result += wc - t2 = time.monotonic_ns() - self.__debug_end("wildcard", start_result, t2 - t1, wc) - return - filter_specifier: list[list[str]] = None - filter_object = tree.children[2] - if filter_object is not None: - if ( # it's an inherited filter from another wildcard - isinstance(filter_object.children[1], lark.Token) - and filter_object.children[1] is not None - and "^" in str(filter_object.children[1]) - ): - filter_wildcard_key = self.__visit(filter_object.children[2], False, True) - filter_specifier = self.__wildcard_filters.get(filter_wildcard_key, None) - if self.__ppp.debug_level == DEBUG_LEVEL.full: - self.__ppp.logger.debug("Filtering choices with inherited filter") - else: - filter_specifier = self.__extract_filter_specifiers(filter_object.children[2]) - if self.__ppp.debug_level == DEBUG_LEVEL.full: - self.__ppp.logger.debug("Filtering choices") - self.__wildcard_filters[wildcard_key] = filter_specifier - if filter_object.children[1] is not None and "#" in str( - filter_object.children[1] - ): # means do not use the filter in this wildcard - if self.__ppp.debug_level == DEBUG_LEVEL.full: - self.__ppp.logger.debug("Ignoring filter") - filter_specifier = None - else: - filter_specifier = self.__ppp.wildcard_obj.get_wildcard_default_filter(wildcard_key) - if filter_specifier is not None: - self.__wildcard_filters[wildcard_key] = filter_specifier - if self.__ppp.debug_level == DEBUG_LEVEL.full: - self.__ppp.logger.debug("Applying default filter") - if ( - len(selected_wildcards) > 1 - and filter_specifier is not None - and any(y[0].isdecimal() for x in filter_specifier for y in x) - ): - self.__ppp.logger.warning( - f"Using a globbing wildcard '{escape_single_quotes(wildcard_key)}' with positional index filters is not recommended!" - ) - var_object = tree.children[3] - variablename = None - variablebackup = None - if var_object is not None: - variablename = str(var_object.children[0]) - variablevalue = self.__visit(var_object.children[1], False, True) - variablebackup = self.__ppp.user_variables.get(variablename, None) - self.__remove_user_variable(variablename) - self.__set_user_variable_value(variablename, variablevalue) - choice_values_all = [] - for wildcard in selected_wildcards: - if wildcard is None: - self.detectedWildcards.append(wc) - self.result += wc - t2 = time.monotonic_ns() - self.__debug_end("wildcard", start_result, t2 - t1, wc) - return - if wildcard.key in self.__seen_wildcards: - self.warn_or_stop( - f"Circular reference detected with wildcard '{escape_single_quotes(self.__seen_wildcards[-1])}' (chain starts at '{escape_single_quotes(self.__seen_wildcards[0])}')!" - ) - continue - self.__seen_wildcards.append(wildcard.key) - if self.__ppp.debug_level == DEBUG_LEVEL.full: - self.__ppp.logger.debug(f"Seen wildcard '{escape_single_quotes(wildcard.key)}'") - (options, choice_values) = self.__check_wildcard_initialization(wildcard) - if options is not None: - if applied_options is None: - applied_options = options - else: - if self.__ppp.debug_level == DEBUG_LEVEL.full: - self.__ppp.logger.debug( - f"Options for wildcard '{escape_single_quotes(wildcard.key)}' are ignored!" - ) - choice_values_all += choice_values - self.result += self.__get_choices(applied_options, choice_values_all, filter_specifier, wildcard_key) - if wildcard_key in self.__wildcard_filters: - del self.__wildcard_filters[wildcard_key] - if variablename is not None: - self.__remove_user_variable(variablename) - if variablebackup is not None: - self.__ppp.user_variables[variablename] = variablebackup - elif self.__ppp.wil_ifwildcards != self.__ppp.IFWILDCARDS_CHOICES.remove: - self.detectedWildcards.append(wc) - self.result += wc - if self.__ppp.debug_level == DEBUG_LEVEL.full: - list_unseen = [f"'{escape_single_quotes(x)}'" for x in self.__seen_wildcards[seen_wildcards_len:]] - self.__ppp.logger.debug(f"Unseen wildcards: {', '.join(list_unseen)}") - self.__seen_wildcards = self.__seen_wildcards[:seen_wildcards_len] - t2 = time.monotonic_ns() - self.__debug_end("wildcard", start_result, t2 - t1, f"'{escape_single_quotes(wc)}'") - - def __extract_filter_specifiers(self, filters: lark.Tree) -> list[list[str]]: - filter_specifier = [] - for or_ in filters.children: - for and_ in or_.children: - label = and_.children[0] - if isinstance(label, lark.Token): - # it's a literal, we can use it directly - filter_specifier.append([str(label)]) - else: - # it's a variable, we need to evaluate it - v = self.__visit(label, False, True) - filter_specifier.append([v]) - return filter_specifier - - def choices(self, tree: lark.Tree): - """ - Process a choices construct in the tree. - """ - t1 = time.monotonic_ns() - start_result = self.result - options = self.__convert_choices_options(tree.children[0], False) - choice_values = [self.__convert_choice(c) for c in tree.children[1::]] - ch = self.__get_original_node_content(tree, "?{...}") - if self.__ppp.wil_process_wildcards: - if self.__ppp.debug_level == DEBUG_LEVEL.full: - self.__ppp.logger.debug("Processing choices:") - self.result += self.__get_choices(options, choice_values) - elif self.__ppp.wil_ifwildcards != self.__ppp.IFWILDCARDS_CHOICES.remove: - self.detectedWildcards.append(ch) - self.result += ch - t2 = time.monotonic_ns() - self.__debug_end("choices", start_result, t2 - t1, f"'{escape_single_quotes(ch)}'") - - def __default__(self, tree): - t1 = time.monotonic_ns() - start_result = self.result - self.__visit(tree.children) - t2 = time.monotonic_ns() - self.__debug_end(tree.data.value, start_result, t2 - t1) - - def start(self, tree): - self.result = "" - t1 = time.monotonic_ns() - self.__visit(tree.children) - attention_processing = self.__ppp.host_config.get("attention", "ok") - # process the found negative tags - for negtag in self.__negtags: - if self.__ppp.cup_mergeattention: - # join consecutive attention elements - for i in range(len(negtag.shell) - 1, 0, -1): - if negtag.shell[i].type == "at" and negtag.shell[i - 1].type == "at": - new_weight = ( # we limit the new weight to two decimals - math.floor(100 * float(negtag.shell[i - 1].data[1]) * float(negtag.shell[i].data[1])) - / 100 - ) - new_weight_str = f"{new_weight:.2f}".rstrip("0").rstrip(".") - if new_weight_str == "0.9" and attention_processing != "parentheses": - new_kind = 1 - elif new_weight_str == "1.1": - new_kind = 2 - else: - new_kind = 3 - negtag.shell[i - 1] = self.AccumulatedShell( - "at", - (new_kind, new_weight_str), - ) - negtag.shell.pop(i) - start = "" - end = "" - for s in negtag.shell: - match s.type: - case "at": - if s.data[0] == 1: - start += "[" - end = "]" + end - elif s.data[0] == 2: - start += "(" - end = ")" + end - else: # 3 - start += "(" - end = f":{s.data[1]})" + end - # case "sc": - case "scb": - start += "[" - end = f"::{s.data}]" + end - case "sca": - start += "[" - end = f":{s.data}]" + end - # case "al": - case "alo": - start += "[" + ("|" * int(s.data["pos"] - 1)) - end = ("|" * int(s.data["len"] - s.data["pos"])) + "]" + end - content = start + negtag.content + end - position = negtag.parameters or "s" - if position.startswith("i"): - n = int(position[1]) - self.insertion_at[n] = [negtag.start, negtag.end] - elif len(content) > 0: - if content not in self.__already_processed: - if self.__ppp.stn_ignore_repeats: - self.__already_processed.append(content) - if self.__ppp.debug_level == DEBUG_LEVEL.full: - self.__ppp.logger.debug( - self.__ppp.format_output(f"Adding content at position {position}: {content}") - ) - if position == "e": - self.add_at["end"].append(content) - elif position.startswith("p"): - n = int(position[1]) - self.add_at["insertion_point"][n].append(content) - else: # position == "s" or invalid - self.add_at["start"].append(content) - else: - self.__ppp.logger.warning(self.__ppp.format_output(f"Ignoring repeated content: {content}")) - t2 = time.monotonic_ns() - self.__debug_end("start", "", t2 - t1) diff --git a/ppp_classes.py b/ppp_classes.py index 6922a34..a3064ac 100644 --- a/ppp_classes.py +++ b/ppp_classes.py @@ -1,10 +1,18 @@ """Pydantic models for the PPP configuration file structure (ppp_config.yaml).""" +from dataclasses import dataclass, field +from logging import Logger import re from enum import Enum -from typing import Literal, Optional +from typing import Any, 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 + + class SUPPORTED_APPS(Enum): comfyui = "comfyui" a1111 = "a1111" @@ -13,6 +21,7 @@ class SUPPORTED_APPS(Enum): sdnext = "sdnext" tests = "tests" # for testing purposes only, not a real app + SUPPORTED_APPS_NAMES = { SUPPORTED_APPS.comfyui: "ComfyUI", SUPPORTED_APPS.sdnext: "SD.Next", @@ -22,6 +31,75 @@ SUPPORTED_APPS_NAMES = { SUPPORTED_APPS.tests: "Tests", } + +class IFWILDCARDS_CHOICES(Enum): + ignore = "ignore" + remove = "remove" + warn = "warn" + stop = "stop" + + +class ONWARNING_CHOICES(Enum): + warn = "warn" + stop = "stop" + + +@dataclass(frozen=True) +class PPPStateOptions: + """Options that can be set for prompt processing.""" + + debug_level: DEBUG_LEVEL = DEBUG_LEVEL.minimal + gen_onwarning: ONWARNING_CHOICES = ONWARNING_CHOICES.warn + wil_process_wildcards: bool = True + wil_keep_choices_order: bool = True + wil_choice_separator: str = ", " + wil_ifwildcards: IFWILDCARDS_CHOICES = IFWILDCARDS_CHOICES.stop + stn_ignore_repeats: bool = True + stn_separator: str = ", " + cup_do_cleanup: bool = True # whether to do cleanup at all (if False, all other cleanup options are ignored) + cup_cleanup_variables: bool = True + cup_extraspaces: bool = True + cup_emptyconstructs: bool = True + cup_extraseparators: bool = True + cup_extraseparators2: bool = True + cup_extraseparators_include_eol: bool = False + cup_breaks: bool = False + cup_breaks_eol: bool = False + cup_ands: bool = False + cup_ands_eol: bool = False + cup_extranetworktags: bool = False + cup_mergeattention: bool = True + rem_removeextranetworktags: bool = False + + +@dataclass(frozen=True) +class PPPState: + """State object passed to various PPP components during prompt processing.""" + + logger: Logger + host_config: dict[str, str] = field(default_factory=dict) + 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) + wildcards_obj: PPPWildcards = field(default_factory=PPPWildcards) + extranetwork_mappings_obj: PPPExtraNetworkMappings = field(default_factory=PPPExtraNetworkMappings) + parsers: dict[str, Lark] = field(default_factory=dict) + + +class PPPInterrupt(Exception): + """ + Custom exception to handle interruptions in the PromptPostProcessor. + This exception can be raised to stop the processing of prompts. + """ + + def __init__(self, message: str = "Processing interrupted.", pos_prefix: str = "", neg_prefix: str = ""): + super().__init__(message) + self.message = message + self.pos_prefix = pos_prefix + self.neg_prefix = neg_prefix + + # ------------------- Host configuration ------------------- AttentionOption = Literal["ok", "parentheses", "disable", "remove", "error"] diff --git a/ppp_comfyui.py b/ppp_comfyui.py index 1d1fda7..364f588 100644 --- a/ppp_comfyui.py +++ b/ppp_comfyui.py @@ -4,7 +4,7 @@ import folder_paths # pylint: disable=import-error # type: ignore import nodes # pylint: disable=import-error # type: ignore from ppp import PromptPostProcessor -from ppp_classes import SUPPORTED_APPS +from ppp_classes import IFWILDCARDS_CHOICES, ONWARNING_CHOICES, SUPPORTED_APPS, PPPStateOptions from ppp_logging import DEBUG_LEVEL, PromptPostProcessorLogFactory from ppp_utils import escape_single_quotes from ppp_wildcards import PPPWildcards @@ -46,9 +46,7 @@ def _resolve_enmappings_folders(override: str = "") -> list[str]: fp3 = None folders_str = ",".join(fp3 or []) if folders_str == "": - folders_str = os.getenv( - "EXTRANETWORKMAPPINGS_DIR", PPPExtraNetworkMappings.DEFAULT_ENMAPPINGS_FOLDER - ) + folders_str = os.getenv("EXTRANETWORKMAPPINGS_DIR", PPPExtraNetworkMappings.DEFAULT_ENMAPPINGS_FOLDER) return [ (f if os.path.isabs(f) else os.path.abspath(os.path.join(folder_paths.models_dir, f))) for f in folders_str.split(",") @@ -64,7 +62,7 @@ class PromptPostProcessorComfyUINode: logger = None def __init__(self): - lf = PromptPostProcessorLogFactory(SUPPORTED_APPS.comfyui) + lf = PromptPostProcessorLogFactory() self.logger = lf.log grammar_filename = os.path.join(os.path.dirname(os.path.realpath(__file__)), "grammar.lark") with open(grammar_filename, "r", encoding="utf-8") as file: @@ -134,7 +132,7 @@ class PromptPostProcessorComfyUINode: }, ), "on_warnings": ( - [e.value for e in PromptPostProcessor.ONWARNING_CHOICES], + [e.value for e in ONWARNING_CHOICES], { "default": PromptPostProcessor.DEFAULT_ONWARNING, "tooltip": "How to handle invalid content warnings", @@ -265,82 +263,78 @@ class PromptPostProcessorComfyUINode: wildcards_folders = _resolve_wildcards_folders(wc_options["wc_wildcards_folders"] if wc_options else "") enmappings_folders = _resolve_enmappings_folders(en_options["en_mappings_folders"] if en_options else "") - options = { - "debug_level": debug_level, - "on_warnings": on_warnings, - "process_wildcards": process_wildcards, - "if_wildcards": ( - wc_options["wc_if_wildcards"] if wc_options else PromptPostProcessor.IFWILDCARDS_CHOICES.stop.value - ), - "choice_separator": ( + options = PPPStateOptions( + debug_level=DEBUG_LEVEL(debug_level), + gen_onwarning=ONWARNING_CHOICES(on_warnings) if on_warnings else PromptPostProcessor.DEFAULT_ONWARNING, + wil_process_wildcards=process_wildcards, + wil_ifwildcards=(wc_options["wc_if_wildcards"] if wc_options else IFWILDCARDS_CHOICES.stop.value), + wil_choice_separator=( wc_options["wc_choice_separator"] if wc_options else PromptPostProcessor.DEFAULT_CHOICE_SEPARATOR ), - "keep_choices_order": ( + wil_keep_choices_order=( wc_options["wc_keep_choices_order"] if wc_options else PromptPostProcessor.DEFAULT_KEEP_CHOICES_ORDER ), - "stn_separator": stn_options["stn_separator"] if stn_options else PromptPostProcessor.DEFAULT_STN_SEPARATOR, - "stn_ignore_repeats": ( + stn_separator=stn_options["stn_separator"] if stn_options else PromptPostProcessor.DEFAULT_STN_SEPARATOR, + stn_ignore_repeats=( stn_options["stn_ignore_repeats"] if stn_options else PromptPostProcessor.DEFAULT_STN_IGNORE_REPEATS ), - "do_cleanup": do_cleanup, - "cleanup_variables": cleanup_variables, - "cleanup_extra_spaces": ( + cup_do_cleanup=do_cleanup, + cup_cleanup_variables=cleanup_variables, + cup_extraspaces=( cup_options["cup_extra_spaces"] if cup_options else PromptPostProcessor.DEFAULT_CUP_EXTRA_SPACES ), - "cleanup_empty_constructs": ( + cup_emptyconstructs=( cup_options["cup_empty_constructs"] if cup_options else PromptPostProcessor.DEFAULT_CUP_EMPTY_CONSTRUCTS ), - "cleanup_extra_separators": ( + cup_extraseparators=( cup_options["cup_extra_separators"] if cup_options else PromptPostProcessor.DEFAULT_CUP_EXTRA_SEPARATORS ), - "cleanup_extra_separators2": ( + cup_extraseparators2=( cup_options["cup_extra_separators2"] if cup_options else PromptPostProcessor.DEFAULT_CUP_EXTRA_SEPARATORS2 ), - "cleanup_extra_separators_include_eol": ( + cup_extraseparators_include_eol=( cup_options["cup_extra_separators_include_eol"] if cup_options else PromptPostProcessor.DEFAULT_CUP_EXTRA_SEPARATORS_INCLUDE_EOL ), - "cleanup_breaks": cup_options["cup_breaks"] if cup_options else PromptPostProcessor.DEFAULT_CUP_BREAKS, - "cleanup_breaks_eol": ( + cup_breaks=cup_options["cup_breaks"] if cup_options else PromptPostProcessor.DEFAULT_CUP_BREAKS, + cup_breaks_eol=( cup_options["cup_breaks_eol"] if cup_options else PromptPostProcessor.DEFAULT_CUP_BREAKS_EOL ), - "cleanup_ands": cup_options["cup_ands"] if cup_options else PromptPostProcessor.DEFAULT_CUP_ANDS, - "cleanup_ands_eol": ( - cup_options["cup_ands_eol"] if cup_options else PromptPostProcessor.DEFAULT_CUP_ANDS_EOL - ), - "cleanup_extranetwork_tags": ( + cup_ands=cup_options["cup_ands"] if cup_options else PromptPostProcessor.DEFAULT_CUP_ANDS, + cup_ands_eol=(cup_options["cup_ands_eol"] if cup_options else PromptPostProcessor.DEFAULT_CUP_ANDS_EOL), + cup_extranetworktags=( cup_options["cup_extranetwork_tags"] if cup_options else PromptPostProcessor.DEFAULT_CUP_EXTRANETWORK_TAGS ), - "cleanup_merge_attention": ( + cup_mergeattention=( cup_options["cup_merge_attention"] if cup_options else PromptPostProcessor.DEFAULT_CUP_MERGE_ATTENTION ), - "remove_extranetwork_tags": ( + rem_removeextranetworktags=( cup_options["cup_remove_extranetwork_tags"] if cup_options else PromptPostProcessor.DEFAULT_CUP_REMOVE_EXTRANETWORK_TAGS ), - } + ) self.wildcards_obj.refresh_wildcards( - debug_level, - wildcards_folders if options["process_wildcards"] else None, + options.debug_level, + wildcards_folders if options.wil_process_wildcards else None, wc_options["wc_wildcards_input"] if wc_options else "", ) self.extranetwork_mappings_obj.refresh_extranetwork_mappings( - debug_level, + options.debug_level, enmappings_folders, en_options["en_mappings_input"] if en_options else "", ) ppp = PromptPostProcessor( self.logger, - self.interrupt, env_info, options, self.grammar_content, + self.interrupt, self.wildcards_obj, self.extranetwork_mappings_obj, ) @@ -386,9 +380,9 @@ class PromptPostProcessorWildcardOptionsComfyUINode: }, ), "if_wildcards": ( - [e.value for e in PromptPostProcessor.IFWILDCARDS_CHOICES], + [e.value for e in IFWILDCARDS_CHOICES], { - "default": PromptPostProcessor.IFWILDCARDS_CHOICES.stop.value, + "default": IFWILDCARDS_CHOICES.stop.value, "tooltip": "How to handle invalid wildcards in the prompt", }, ), @@ -761,17 +755,31 @@ class PromptPostProcessorWildcardConcatComfyUINode: """ NONE_OPTION = "(none)" - _wildcards_obj = None - _wildcards_log = None + _ppp = None + # _state = None + # _tree = None + _wildcard_key_map: dict[str, str] = {} @classmethod def _ensure_wildcards(cls): - if cls._wildcards_obj is None: - lf = PromptPostProcessorLogFactory(SUPPORTED_APPS.comfyui) - cls._wildcards_log = lf.log - cls._wildcards_obj = PPPWildcards(cls._wildcards_log) - cls._wildcards_obj.refresh_wildcards(DEBUG_LEVEL.minimal, _resolve_wildcards_folders()) - return cls._wildcards_obj + if cls._ppp is None: + lf = PromptPostProcessorLogFactory() + cls._ppp = PromptPostProcessor( + lf.log, + { + "app": SUPPORTED_APPS.comfyui.value, + "models_path": folder_paths.models_dir, + "model_filename": "", + "model_class": "", + "property_base": None, + }, + PPPStateOptions(debug_level=DEBUG_LEVEL.minimal), + wildcards_obj=PPPWildcards(lf.log), + ) + # We need to populate the wildcard options in the tree to get the descriptions for the dropdown labels + cls._ppp.state.wildcards_obj.refresh_wildcards(cls._ppp.state.options.debug_level, _resolve_wildcards_folders()) + cls._ppp.init_wildcards_options() + return cls._ppp.state.wildcards_obj @classmethod def get_wildcard_keys(cls, filter_prefix=""): @@ -779,7 +787,16 @@ class PromptPostProcessorWildcardConcatComfyUINode: keys = sorted(wc_obj.wildcards.keys()) if filter_prefix: keys = [k for k in keys if k.startswith(filter_prefix)] - return [cls.NONE_OPTION] + keys + cls._wildcard_key_map = {cls.NONE_OPTION: cls.NONE_OPTION} + labels = [] + for k in keys: + wc = wc_obj.wildcards[k] + pretty_key = k.replace("_", " ") + desc = str(wc.options["description"]) if wc.options and "description" in wc.options else None + label = f"{pretty_key} ({desc})" if desc else pretty_key + labels.append(label) + cls._wildcard_key_map[label] = k + return [cls.NONE_OPTION] + labels @classmethod def INPUT_TYPES(cls): @@ -849,8 +866,9 @@ class PromptPostProcessorWildcardConcatComfyUINode: wildcard_10, previous_prompt=None, ): + key_map = self.__class__._wildcard_key_map # pylint: disable=protected-access parts = [ - f"__{w}__" + f"__{key_map.get(w, w)}__" for w in [ wildcard_1, wildcard_2, diff --git a/ppp_common.py b/ppp_common.py new file mode 100644 index 0000000..0eefe88 --- /dev/null +++ b/ppp_common.py @@ -0,0 +1,193 @@ +import ast +import logging +import os +import re +import textwrap +import time +import lark + +from ppp_logging import DEBUG_LEVEL +from ppp_classes import ONWARNING_CHOICES, PPPInterrupt, PPPState +from ppp_utils import escape_single_quotes, format_output + + +def parse_prompt( + state: PPPState, + prompt_description: str, + prompt: str, + parser: lark.Lark, + raise_parsing_error: bool = False, +): + """ + Parses a prompt using the specified parser. + + Args: + prompt_description (str): The description of the prompt. + prompt (str): The prompt to be parsed. + parser (lark.Lark): The parser to be used. + raise_parsing_error (bool): Whether to raise a parsing error. + + Returns: + Tree: The parsed prompt. + """ + t1 = time.monotonic_ns() + parsed_prompt = None + try: + if state.options.debug_level == DEBUG_LEVEL.full: + state.logger.debug( + format_output( + f"Parsing {prompt_description}: '{escape_single_quotes(prompt)}'" + ) + ) + parsed_prompt = parser.parse(prompt) + # we store the contents so we can use them later even if the meta position is not valid anymore + if isinstance(parsed_prompt, lark.Tree): + for n in parsed_prompt.iter_subtrees(): + if isinstance(n, lark.Tree): + if n.meta.empty: + n.meta.content = "" + else: + n.meta.content = prompt[n.meta.start_pos : n.meta.end_pos] + except lark.exceptions.UnexpectedInput: + if raise_parsing_error: + raise + state.logger.exception( + format_output(f"Parsing failed on prompt!: {escape_single_quotes(prompt)}") + ) + t2 = time.monotonic_ns() + if state.options.debug_level == DEBUG_LEVEL.full: + state.logger.debug( + f"Parse {prompt_description} time: {(t2 - t1) / 1_000_000_000:.3f} seconds" + ) + if parsed_prompt: + state.logger.debug( + "Tree:\n" + + textwrap.indent( + re.sub( + r"\n$", + "", + ( + parsed_prompt.pretty() + if isinstance(parsed_prompt, lark.Tree) + else parsed_prompt + ), + ), + " ", + ) + ) + return parsed_prompt + + +def warn_or_stop(state: PPPState, is_negative: bool, message: str, e: Exception = None): + INVALID_CONTENT_STOP = "INVALID CONTENT! {0}\nBREAK " + if state.options.gen_onwarning == ONWARNING_CHOICES.stop: + raise PPPInterrupt( + message, + INVALID_CONTENT_STOP.format(message) if not is_negative else "", + INVALID_CONTENT_STOP.format(message) if is_negative else "", + ) from e + state.logger.warning(format_output(message)) + +def load_grammar() -> str: + # Process with lark (debug with https://www.lark-parser.org/ide/) + grammar_filename = os.path.join(os.path.dirname(os.path.realpath(__file__)), "grammar.lark") + with open(grammar_filename, "r", encoding="utf-8") as file: + grammar_content = file.read() + return grammar_content + + +def preprocess_grammar(grammar_content: str, options: dict[str, bool], logger: logging.Logger = None) -> str: + """ + Preprocesses the grammar content to handle conditional compilation directives. + + Args: + grammar_content (str): The raw grammar content. + options (dict[str,bool]): Options for preprocessing. + + Returns: + str: The preprocessed grammar content. + """ + lines = grammar_content.split("\n") + result_lines = [] + skip_current_block = [] + all_blocks_skipped = [] + + def eval_bool_expr(expr: str, constants: dict[str, bool]) -> bool: + """ + Evaluates a boolean expression using known constants. + Supports: and, or, not, parentheses, and named constants. + + Args: + expr (str): The boolean expression to evaluate. + constants (dict[str, bool]): A dictionary of constant values. + Returns: + bool: The result of the evaluated expression. + """ + tree = ast.parse(expr, mode="eval") + + def _eval(node) -> bool: + if isinstance(node, ast.Expression): + return _eval(node.body) + if isinstance(node, ast.BoolOp): + if isinstance(node.op, ast.And): + return all(_eval(v) for v in node.values) + if isinstance(node.op, ast.Or): + return any(_eval(v) for v in node.values) + if isinstance(node, ast.UnaryOp) and isinstance(node.op, ast.Not): + return not _eval(node.operand) + if isinstance(node, ast.Name): + return bool(constants[node.id]) # raises KeyError for unknown names + if isinstance(node, ast.Constant) and isinstance(node.value, bool): + return node.value + raise ValueError(f"Unsupported construct: {ast.dump(node)}") + + return _eval(tree) + + for line in lines: + stripped_line = line.strip() + if stripped_line.startswith("//#if"): + # Extract condition from the #if directive + conditions = stripped_line[5:].strip() + # Evaluate the conditions + skip_current_block.append(not eval_bool_expr(conditions, options)) + all_blocks_skipped.append(skip_current_block[-1]) + elif stripped_line.startswith("//#elif"): + if not skip_current_block: + if logger: + logger.warning("Unmatched //#elif directive found in grammar content.") + elif all_blocks_skipped[-1]: + # Extract condition from the #elif directive + conditions = stripped_line[7:].strip() + # Evaluate the conditions + skip_current_block[-1] = not eval_bool_expr(conditions, options) + if not skip_current_block[-1]: + all_blocks_skipped[-1] = False + else: + skip_current_block[-1] = True + elif stripped_line.startswith("//#else"): + if not skip_current_block: + if logger: + logger.warning("Unmatched //#else directive found in grammar content.") + elif all_blocks_skipped[-1]: + skip_current_block[-1] = False + else: + skip_current_block[-1] = True + elif stripped_line.startswith("//#endif"): + if not skip_current_block: + if logger: + logger.warning("Unmatched //#endif directive found in grammar content.") + else: + skip_current_block.pop() + all_blocks_skipped.pop() + elif stripped_line.startswith("//#"): + if logger: + logger.warning(f"Unrecognized directive found in grammar content: {stripped_line}") + elif not any(skip_current_block): + # Include the line if we're not skipping any current block + result_lines.append(stripped_line) + # Check for unclosed blocks at the end + if skip_current_block: + raise PPPInterrupt( + f"Found {len(skip_current_block)} unclosed conditional directive(s) at the end of the grammar file" + ) + return "\n".join(result_lines) diff --git a/ppp_enmappings.py b/ppp_enmappings.py index c353f35..efd7ae3 100644 --- a/ppp_enmappings.py +++ b/ppp_enmappings.py @@ -68,7 +68,7 @@ class PPPExtraNetworkMappings: DEFAULT_ENMAPPINGS_FOLDER = "extranetworkmappings" LOCALINPUT_FILENAME = R"//INPUT\\" - def __init__(self, logger): + def __init__(self, logger=None): self.__logger: logging.Logger = logger self.__debug_level = DEBUG_LEVEL.none self.__enmappings_folders = [] @@ -95,7 +95,7 @@ class PPPExtraNetworkMappings: """ self.__debug_level = debug_level self.__enmappings_folders = enmappings_folders or [] - # if self.__debug_level != DEBUG_LEVEL.none: + # if self.__debug_level != DEBUG_LEVEL.none and self.__logger: # self.__logger.info("Refreshing extra network mappings...") # t1 = time.monotonic_ns() self.cached_mappings = {} @@ -118,7 +118,7 @@ class PPPExtraNetworkMappings: self.extranetwork_mappings = {} self.__enmappings_files = {} # t2 = time.monotonic_ns() - # if self.__debug_level != DEBUG_LEVEL.none: + # if self.__debug_level != DEBUG_LEVEL.none and self.__logger: # self.__logger.info(f"Extra network mappings refresh time: {(t2 - t1) / 1_000_000_000:.3f} seconds") # def get_extranetwork_mappings(self, key: str) -> list[PPPENMapping]: @@ -143,7 +143,7 @@ class PPPExtraNetworkMappings: debug (bool): Whether to print debug messages or not. """ last_modified_cached = self.__enmappings_files.get(full_path, None) # a time or a hash - if debug and last_modified_cached is not None and self.__debug_level != DEBUG_LEVEL.none: + if debug and last_modified_cached is not None and self.__debug_level != DEBUG_LEVEL.none and self.__logger: if full_path == self.LOCALINPUT_FILENAME: self.__logger.debug("Removing extra network mappings from input") else: @@ -170,7 +170,7 @@ class PPPExtraNetworkMappings: if extension not in (".yaml", ".yml", ".json"): return self.__remove_extranetwork_mappings_from_path(full_path, False) - if last_modified_cached is not None and self.__debug_level != DEBUG_LEVEL.none: + if last_modified_cached is not None and self.__debug_level != DEBUG_LEVEL.none and self.__logger: self.__logger.debug(f"Updating extra network mappings from file: {full_path}") self.__get_extranetwork_mappings_in_structured_file(full_path) self.__enmappings_files[full_path] = last_modified @@ -187,14 +187,15 @@ class PPPExtraNetworkMappings: if h == new_h: return self.__remove_extranetwork_mappings_from_path(self.LOCALINPUT_FILENAME, False) - if h is not None and self.__debug_level != DEBUG_LEVEL.none: + if h is not None and self.__debug_level != DEBUG_LEVEL.none and self.__logger: self.__logger.debug("Updating extra network mappings from input") enmappings_input = enmappings_input.strip() if enmappings_input != "": try: content = yaml.safe_load(enmappings_input) except yaml.YAMLError as e: - self.__logger.warning(f"Invalid format for input extra network mappings: {e}") + if self.__logger: + self.__logger.warning(f"Invalid format for input extra network mappings: {e}") return if content is not None: self.__add_extranetwork_mapping(content, self.LOCALINPUT_FILENAME) @@ -209,28 +210,33 @@ class PPPExtraNetworkMappings: full_path (str): The path to the file that contains it. """ if not isinstance(content, dict): - self.__logger.warning(f"Invalid extra network mapping in file '{escape_single_quotes(full_path)}'!") + if self.__logger: + self.__logger.warning(f"Invalid extra network mapping in file '{escape_single_quotes(full_path)}'!") return for kind, maps in content.items(): if not isinstance(maps, dict): - self.__logger.warning( - f"Invalid extra network mapping definition for '{escape_single_quotes(kind)}:*' in file '{escape_single_quotes(full_path)}'!" - ) + if self.__logger: + self.__logger.warning( + f"Invalid extra network mapping definition for '{escape_single_quotes(kind)}:*' in file '{escape_single_quotes(full_path)}'!" + ) else: for name, variants in maps.items(): key = f"{kind}:{name}" if not isinstance(variants, list): - self.__logger.warning( - f"Invalid extra network mapping definition for '{escape_single_quotes(key)}' in file '{escape_single_quotes(full_path)}'!" - ) + if self.__logger: + self.__logger.warning( + f"Invalid extra network mapping definition for '{escape_single_quotes(key)}' in file '{escape_single_quotes(full_path)}'!" + ) elif self.extranetwork_mappings.get(key, None) is not None: - self.__logger.warning( - f"Duplicate extra network mapping '{escape_single_quotes(key)}' in file '{escape_single_quotes(full_path)}' and '{escape_single_quotes(self.extranetwork_mappings[key].file)}'!" - ) + if self.__logger: + self.__logger.warning( + f"Duplicate extra network mapping '{escape_single_quotes(key)}' in file '{escape_single_quotes(full_path)}' and '{escape_single_quotes(self.extranetwork_mappings[key].file)}'!" + ) elif not isinstance(variants, list) or not all(isinstance(v, dict) for v in variants): - self.__logger.warning( - f"Invalid extra network mapping definition for '{escape_single_quotes(key)}' in file '{escape_single_quotes(full_path)}'!" - ) + if self.__logger: + self.__logger.warning( + f"Invalid extra network mapping definition for '{escape_single_quotes(key)}' in file '{escape_single_quotes(full_path)}'!" + ) else: self.extranetwork_mappings[key] = PPPENMapping(full_path, kind, name, variants) @@ -247,16 +253,18 @@ class PPPExtraNetworkMappings: with open(full_path, "r", encoding="utf-8") as file: content = yaml.safe_load(file) except: # pylint: disable=bare-except - self.__logger.warning( - f"Could not read file '{escape_single_quotes(full_path)}' with utf-8 encoding, trying windows-1252..." - ) + if self.__logger: + self.__logger.warning( + f"Could not read file '{escape_single_quotes(full_path)}' with utf-8 encoding, trying windows-1252..." + ) with open(full_path, "r", encoding="windows-1252") as file: content = yaml.safe_load(file) self.__add_extranetwork_mapping(content, full_path) except Exception as e: # pylint: disable=broad-except - self.__logger.error( - f"Error reading extra network mappings from file '{escape_single_quotes(full_path)}': {e}" - ) + if self.__logger: + self.__logger.error( + f"Error reading extra network mappings from file '{escape_single_quotes(full_path)}': {e}" + ) def __get_extranetwork_mappings_in_directory(self, directory: str): """ @@ -266,9 +274,10 @@ class PPPExtraNetworkMappings: directory (str): The path to the directory. """ if not os.path.exists(directory): - self.__logger.warning( - f"Extra network mappings directory '{escape_single_quotes(directory)}' does not exist!" - ) + if self.__logger: + self.__logger.warning( + f"Extra network mappings directory '{escape_single_quotes(directory)}' does not exist!" + ) return for filename in os.listdir(directory): full_path = os.path.abspath(os.path.join(directory, filename)) diff --git a/ppp_logging.py b/ppp_logging.py index 113f578..1574d1e 100644 --- a/ppp_logging.py +++ b/ppp_logging.py @@ -2,7 +2,6 @@ from enum import Enum import logging import sys import copy -from ppp_classes import SUPPORTED_APPS class DEBUG_LEVEL(Enum): @@ -53,7 +52,7 @@ class PromptPostProcessorLogFactory: # pylint: disable=too-few-public-methods colored_record.levelname = f"{seq}{levelname:8s}{self.COLORS['RESET']}" return super().format(colored_record) - def __init__(self, app: SUPPORTED_APPS = None, filename = None): # pylint: disable=unused-argument + def __init__(self, filename = None): """ Initializes the PromptPostProcessor class. diff --git a/ppp_tree.py b/ppp_tree.py new file mode 100644 index 0000000..5b8265f --- /dev/null +++ b/ppp_tree.py @@ -0,0 +1,1630 @@ +from collections import namedtuple +import math +import re +import textwrap +import time +from typing import Optional +import lark +import numpy as np + +from ppp_classes import IFWILDCARDS_CHOICES, PPPState +from ppp_enmappings import PPPENMappingVariant +from ppp_logging import DEBUG_LEVEL +from ppp_utils import escape_single_quotes, format_output +from ppp_common import parse_prompt, warn_or_stop +from ppp_wildcards import PPPWildcard + + +class TreeProcessor(lark.visitors.Interpreter): + """ + A class for interpreting and processing a tree generated by the prompt parser. + + Args: + state (PPPState): The state object containing the current processing state. + rng (numpy.random.Generator): The random number generator. + + Attributes: + add_at (dict): The dictionary to store the content to be added at different positions of the negative prompt. + insertion_at (list): The list of insertion points in the negative prompt. + detectedWildcards (list): The list of detected invalid wildcards or choices. + result (str): The final processed prompt. + """ + + def __init__(self, state: PPPState, rng: np.random.Generator): + super().__init__() + self.__state = state + self.__debug_level = state.options.debug_level + self.__rng = rng + self.AccumulatedShell = namedtuple("AccumulatedShell", ["type", "data"]) + self.NegTag = namedtuple("NegTag", ["start", "end", "content", "parameters", "shell"]) + self.__shell: list[self.AccumulatedShell] = [] # type: ignore + self.__negtags: list[self.NegTag] = [] # type: ignore + self.__already_processed: list[str] = [] + self.__is_negative = False + self.__wildcard_filters = {} + self.__seen_wildcards: list[str] = [] + self.add_at: dict = {"start": [], "insertion_point": [[] for x in range(10)], "end": []} + self.insertion_at: list[tuple[int, int]] = [None for x in range(10)] + self.detectedWildcards: list[str] = [] + self.result = "" + + def warn_or_stop(self, message: str, e: Exception = None): + warn_or_stop(self.__state, self.__is_negative, message, e) + + def start_visit(self, prompt_description: str, parsed_prompt: lark.Tree, is_negative: bool = False) -> str: + """ + Start the visit process. + + Args: + prompt_description (str): The description of the prompt. + parsed_prompt (Tree): The parsed prompt. + is_negative (bool): Whether the prompt is negative or not. + + Returns: + str: The processed prompt. + """ + t1 = time.monotonic_ns() + self.__is_negative = is_negative + if self.__debug_level != DEBUG_LEVEL.none: + self.__state.logger.info(f"Processing {prompt_description}...") + self.visit(parsed_prompt) + t2 = time.monotonic_ns() + if self.__debug_level != DEBUG_LEVEL.none: + self.__state.logger.info(f"Process {prompt_description} time: {(t2 - t1) / 1_000_000_000:.3f} seconds") + return self.result + + def __visit( + self, + node: lark.Tree | lark.Token | list[lark.Tree | lark.Token] | None, + restore_state: bool = False, + discard_content: bool = False, + ) -> str: + """ + Visit a node in the tree and process it or accumulate its value if it is a Token. + + Args: + node (Tree|Token|list): The node or list of nodes to visit. + restore_state (bool): Whether to restore the state after visiting the node. + discard_content (bool): Whether to discard the content of the node. + + Returns: + str: The result of the visit. + """ + backup_result = self.result + # if self.__debug_level == DEBUG_LEVEL.full: + # self.__state.logger.debug(f"Visiting node {node}.") + if restore_state: + # if self.__debug_level == DEBUG_LEVEL.full: + # self.__state.logger.debug("Backing up state before visiting.") + backup_shell = self.__shell.copy() + backup_negtags = self.__negtags.copy() + backup_already_processed = self.__already_processed.copy() + 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() + if node is not None: + if isinstance(node, list): + for child in node: + self.__visit(child) + elif isinstance(node, lark.Tree): + self.visit(node) + elif isinstance(node, lark.Token): + self.result += node + len_backup = len(backup_result) + # if self.result[:len_backup] == backup_result: # this is only necessary if we call parse_prompt with a parser from "start", because it resets the result + added_result = self.result[len_backup:] + # else: + # added_result = self.result + if discard_content or restore_state: + self.result = backup_result + if restore_state: + # if self.__debug_level == DEBUG_LEVEL.full: + # self.__state.logger.debug("Restoring state after visiting.") + self.__shell = backup_shell + self.__negtags = backup_negtags + self.__already_processed = backup_already_processed + 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) + return added_result + + def __get_original_node_content(self, node: lark.Tree | lark.Token, default=None) -> str: + """ + Get the original content of a node. + + Args: + node (Tree|Token): The node to get the content from. + default: The default value to return if the content is not found. + + Returns: + str: The original content of the node. + """ + return node.meta.content if hasattr(node, "meta") and node.meta is not None and not node.meta.empty else default + + def __get_user_variable_value(self, name: str, evaluate=True, visit=False) -> str: + """ + Get the value of a user variable. + + Args: + name (str): The name of the user variable. + evaluate (bool): Whether to evaluate the variable. + visit (bool): Whether to also visit the variable (add to result). + + Returns: + str: The value of the user variable. + """ + v = self.__state.user_variables.get(name, None) + if v is not None: + visited = False + if isinstance(v, lark.Tree): + if evaluate: + v = self.__visit(v, not visit) + visited = visit + else: + v = self.__get_original_node_content(v, "not evaluated yet") + if visit and not visited: + self.result += v + return v + + def get_final_user_variable(self, name: str) -> str: + return self.__get_user_variable_value(name, True, False) + + def __set_user_variable_value(self, name: str, value: str): + """ + Set the value of a user variable. + + Args: + name (str): The name of the user variable. + value (str): 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. + + Args: + construct (str): The name of the construct. + start_result (str): The initial result. + duration (int): The duration of the processing in ns. + info: Additional information to log. + """ + if self.__debug_level == DEBUG_LEVEL.full: + info = f"({info}) " if info is not None and info != "" else "" + output = self.result[len(start_result) :] + if output != "": + output = f" >> '{escape_single_quotes(output)}'" + self.__state.logger.debug( + format_output(f"TreeProcessor.{construct} {info}({duration / 1_000_000_000:.3f} seconds){output}") + ) + + def __resolve_cond_value(self, c: str): + """Resolve a condition value: try int first, fall back to variable lookup.""" + try: + return int(c) + except ValueError: + # Bare identifier - resolve as variable reference + if c.startswith("_"): + val = self.__state.system_variables.get(c, None) + if val is None: + val = "" + self.warn_or_stop(f"Unknown system variable {c}") + else: + val = self.__get_user_variable_value(c) + if val is None: + val = "" + self.warn_or_stop(f"Unknown user variable {c}") + return val.lower() if isinstance(val, str) else val + + def __eval_basiccondition(self, cond_var: str, cond_comp: str, cond_value: str | list[str]) -> bool: + """ + Evaluate a condition based on the given variable, comparison, and value. + + Args: + cond_var (str): The variable to be compared. + cond_comp (str): The comparison operator. + cond_value (str or list[str]): The value to be compared with. + + Returns: + bool: The result of the condition evaluation. + """ + if cond_var.lower() == "false": + var_value = "false" + elif cond_var.lower() == "true": + var_value = "true" + elif cond_var.startswith("_"): # system variable + var_value = self.__state.system_variables.get(cond_var, None) + if var_value is None: + var_value = "" + self.warn_or_stop(f"Unknown system variable {cond_var}") + else: # user variable + var_value = self.__get_user_variable_value(cond_var) + if var_value is None: + var_value = "" + self.warn_or_stop(f"Unknown user variable {cond_var}") + if isinstance(var_value, str): + var_value = var_value.lower() + if isinstance(cond_value, list): + comp_ops = { + "contains": lambda x, y: y in x, + "in": lambda x, y: x == y, + } + else: + cond_value = [cond_value] + comp_ops = { + "eq": lambda x, y: x == y, + "ne": lambda x, y: x != y, + "gt": lambda x, y: x > y, + "lt": lambda x, y: x < y, + "ge": lambda x, y: x >= y, + "le": lambda x, y: x <= y, + "contains": lambda x, y: y in x, + "truthy": lambda x, y: bool(x), + } + if cond_comp not in comp_ops: + return False + cond_value_adjusted = list( + ( + c[1:-1].lower() + if c.startswith('"') or c.startswith("'") + else ( + True + if c.lower() == "true" + else False if c.lower() == "false" or c == "" else self.__resolve_cond_value(c) + ) + ) + for c in cond_value + ) + result = False + for c in cond_value_adjusted: + if isinstance(c, str): + var_value_adjusted = var_value + elif isinstance(c, bool) and var_value != "false" and var_value != "" and var_value is not False: + var_value_adjusted = True + elif isinstance(c, bool) and (var_value != "true" or var_value is False): + var_value_adjusted = False + else: + try: + var_value_adjusted = int(var_value) + except (ValueError, TypeError): + self.warn_or_stop( + f"Cannot convert variable value '{escape_single_quotes(var_value)}' to integer for comparison" + ) + return False + result = comp_ops[cond_comp](var_value_adjusted, c) + if result: + break + return result + + def __eval_condition(self, condition: lark.Tree) -> bool: + """ + Evaluate an if condition based on the given condition tree. + + Args: + condition (Node): The condition tree to be evaluated. + + Returns: + bool: The result of the if condition evaluation. + """ + # self.__state.logger.debug(f"__eval_condition {condition.data}") + if condition.data == "operation_and": + cond_result = True + for c in condition.children: + cond_result = cond_result and self.__eval_condition(c) + if not cond_result: + break + elif condition.data == "operation_or": + cond_result = False + for c in condition.children: + cond_result = cond_result or self.__eval_condition(c) + if cond_result: + break + elif condition.data == "operation_not": + cond_result = not self.__eval_condition(condition.children[0]) + else: # truthy_operand / comparison_simple_value / comparison_list_value + # we get the name of the variable + cond_var = str(condition.children[0]) + poscomp = 1 + invert = False + if poscomp >= len(condition.children): + # no condition, just a variable + cond_comp = "truthy" + cond_value = "true" + else: + # we get the comparison (with possible not) and the value + cond_comp = str(condition.children[poscomp]) + if cond_comp == "not": + invert = not invert + poscomp += 1 + cond_comp = str(condition.children[poscomp]) + poscomp += 1 + cond_value_node = condition.children[poscomp] + cond_value = ( + list(str(v) for v in cond_value_node.children) + if isinstance(cond_value_node, (lark.Tree, list)) + else str(cond_value_node) if isinstance(cond_value_node, lark.Token) else cond_value_node + ) + cond_result = self.__eval_basiccondition(cond_var, cond_comp, cond_value) + if invert: + cond_result = not cond_result + return cond_result + + def promptcomp(self, tree: lark.Tree): + """ + Process a prompt composition construct in the tree. + """ + start_result = self.result + t1 = time.monotonic_ns() + self.__visit(tree.children[0]) + and_processing = self.__state.host_config.get("and", "ok") + if len(tree.children) > 1: + and_replacements = { + "eol": ("replaced with EOL", "\n"), + "comma": ("replaced with COMMA", ", "), + "remove": ("removed", " "), + } + if tree.children[1] is not None: + self.result += f":{tree.children[1]}" + for i in range(2, len(tree.children), 3): + if and_processing in and_replacements.keys(): + self.result = ( + self.result.rstrip() + + and_replacements[and_processing][1] + + self.__visit(tree.children[i + 1], False, True).lstrip() + ) + if self.__debug_level == DEBUG_LEVEL.full: + self.__state.logger.debug(f"AND construct {and_replacements[and_processing][0]}") + elif and_processing == "error": + self.warn_or_stop("AND constructs are not allowed!") + else: # and_processing == "ok": + if self.__state.options.cup_ands: + self.result = re.sub(r"[, ]+$", "\n" if self.__state.options.cup_ands_eol else " ", self.result) + if self.result[-1:].isalnum(): # add space if needed + self.result += " " + self.result += "AND" + added_result = self.__visit(tree.children[i + 1], False, True) + if self.__state.options.cup_ands: + added_result = re.sub(r"^[, ]+", " ", added_result) + if added_result[0:1].isalnum(): # add space if needed + added_result = " " + added_result + self.result += added_result + if tree.children[i + 2] is not None: + self.result += f":{tree.children[i+2]}" + t2 = time.monotonic_ns() + self.__debug_end("promptcomp", start_result, t2 - t1) + + def scheduled(self, tree: lark.Tree): + """ + Process a scheduling construct in the tree and add it to the accumulated shell. + """ + start_result = self.result + t1 = time.monotonic_ns() + before = tree.children[0] + after = tree.children[-2] + pos_str = tree.children[-1] + pos = float(pos_str) + if pos >= 1: + pos = int(pos) + scheduling_processing = self.__state.host_config.get("scheduling", "ok") + if scheduling_processing == "before": + if self.__debug_level == DEBUG_LEVEL.full: + self.__state.logger.debug("Scheduling construct removed, taking before option") + if before is not None: + self.__visit(before) + elif scheduling_processing == "after": + if self.__debug_level == DEBUG_LEVEL.full: + self.__state.logger.debug("Scheduling construct removed, taking after option") + if after is not None: + self.__visit(after) + elif scheduling_processing == "first": + if self.__debug_level == DEBUG_LEVEL.full: + self.__state.logger.debug("Scheduling construct removed, taking first option") + if before is not None: + self.__visit(before) + elif after is not None: + self.__visit(after) + elif scheduling_processing == "remove": + if self.__debug_level == DEBUG_LEVEL.full: + self.__state.logger.debug("Scheduling construct removed") + elif scheduling_processing == "error": + self.warn_or_stop("Scheduling constructs are not allowed!") + else: # scheduling_processing == "ok" + # self.__shell.append(self.AccumulatedShell("sc", pos)) + self.result += "[" + if before is not None: + if self.__debug_level == DEBUG_LEVEL.full: + self.__state.logger.debug(f"Shell scheduled before with position {pos}") + self.__shell.append(self.AccumulatedShell("scb", pos)) + self.__visit(before) + self.__shell.pop() + if self.__debug_level == DEBUG_LEVEL.full: + self.__state.logger.debug(f"Shell scheduled after with position {pos}") + self.__shell.append(self.AccumulatedShell("sca", pos)) + self.result += ":" + self.__visit(after) + self.__shell.pop() + if self.__state.options.cup_emptyconstructs and re.fullmatch( + re.escape(start_result) + r"\[:\s*", self.result + ): + self.result = start_result + else: + self.result += f":{pos_str}]" + # self.__shell.pop() + t2 = time.monotonic_ns() + self.__debug_end("scheduled", start_result, t2 - t1, pos_str) + + def alternate(self, tree: lark.Tree): + """ + Process an alternation construct in the tree and add it to the accumulated shell. + """ + start_result = self.result + t1 = time.monotonic_ns() + alternation_processing = self.__state.host_config.get("alternation", "ok") + if alternation_processing == "first": + if self.__debug_level == DEBUG_LEVEL.full: + self.__state.logger.debug("Alternation construct removed, taking first option") + self.__visit(tree.children[0]) + elif alternation_processing == "remove": + if self.__debug_level == DEBUG_LEVEL.full: + self.__state.logger.debug("Alternation construct removed") + elif alternation_processing == "error": + self.warn_or_stop("Alternation constructs are not allowed!") + else: # alternation_processing == "ok" + # self.__shell.append(self.AccumulatedShell("al", len(tree.children))) + self.result += "[" + for i, opt in enumerate(tree.children): + if self.__debug_level == DEBUG_LEVEL.full: + self.__state.logger.debug(f"Shell alternate option {i+1}") + self.__shell.append(self.AccumulatedShell("alo", {"pos": i + 1, "len": len(tree.children)})) + if i > 0: + self.result += "|" + self.__visit(opt) + self.__shell.pop() + self.result += "]" + if self.__state.options.cup_emptyconstructs and re.fullmatch( + re.escape(start_result) + r"\[\s*\]", self.result + ): + self.result = start_result + # self.__shell.pop() + t2 = time.monotonic_ns() + self.__debug_end("alternate", start_result, t2 - t1) + + def attention(self, tree: lark.Tree): + """ + Process a attention change construct in the tree and add it to the accumulated shell. + """ + start_result = self.result + t1 = time.monotonic_ns() + # weight_kind: -1: remove, 0=none, 1=decrease, 2=increase, 3=specific + if len(tree.children) == 2: + weight_str = tree.children[-1] + if weight_str is not None: + weight_kind = 3 # specific weight + weight = float(weight_str) + else: + weight_kind = 2 # increase attention + weight = 1.1 + weight_str = "1.1" + else: + weight_kind = 1 # decrease attention + weight = 0.9 + weight_str = "0.9" + if self.__debug_level == DEBUG_LEVEL.full: + self.__state.logger.debug(f"Shell attention with weight {weight}") + current_tree = tree.children[0] + if self.__state.options.cup_mergeattention: + while isinstance(current_tree, lark.Tree) and current_tree.data == "attention": + # we merge the weights + if len(current_tree.children) == 2: + inner_weight = current_tree.children[-1] + if inner_weight is not None: + inner_weight = float(inner_weight) + else: + inner_weight = 1.1 + else: + inner_weight = 0.9 + weight *= inner_weight + current_tree = current_tree.children[0] + weight = math.floor(weight * 100) / 100 # we round to 2 decimals + weight_str = f"{weight:.2f}".rstrip("0").rstrip(".") + if weight_str == "0.9": + weight_kind = 1 + elif weight_str == "1.1": + weight_kind = 2 + else: + weight_kind = 3 + attention_processing = self.__state.host_config.get("attention", "ok") + if attention_processing == "parentheses": + if weight_kind == 1: + weight_kind = 3 + weight_str = "0.9" + if self.__debug_level == DEBUG_LEVEL.full: + self.__state.logger.debug("Converted to parentheses format") + elif attention_processing == "disable": + weight_kind = 0 + if self.__debug_level == DEBUG_LEVEL.full: + self.__state.logger.debug("Attention construct disabled") + elif attention_processing == "remove": + weight_kind = -1 + if self.__debug_level == DEBUG_LEVEL.full: + self.__state.logger.debug("Attention construct removed") + elif attention_processing == "error": + self.warn_or_stop("Attention constructs are not allowed!") + # else: attention_processing == "ok": + if weight_kind == -1: + # we just ignore the attention construct + pass + elif weight_kind == 0: + # we just visit the content without adding any attention + self.__visit(current_tree) + else: + self.__shell.append(self.AccumulatedShell("at", (weight_kind, weight_str))) + if weight_kind == 1: + starttag = "[" + self.result += starttag + self.__visit(current_tree) + endtag = "]" + elif weight_kind == 2: + starttag = "(" + self.result += starttag + self.__visit(current_tree) + endtag = ")" + else: # weight_kind == 3 + starttag = "(" + self.result += starttag + self.__visit(current_tree) + endtag = f":{weight_str})" + if self.__state.options.cup_emptyconstructs and re.fullmatch( + re.escape(start_result + starttag) + r"\s*", self.result + ): + self.result = start_result + else: + self.result += endtag + self.__shell.pop() + t2 = time.monotonic_ns() + self.__debug_end("attention", start_result, t2 - t1, weight_str) + + def commandstn(self, tree: lark.Tree): + """ + Process a send to negative command in the tree and add it to the list of negative tags. + """ + start_result = self.result + info = None + t1 = time.monotonic_ns() + if not self.__is_negative: + negtagparameters = tree.children[0] + if negtagparameters is not None: + parameters = str(negtagparameters) + else: + parameters = "" + content = self.__visit(tree.children[1::], False, True) + self.__negtags.append( + self.NegTag(len(self.result), len(self.result), content, parameters, self.__shell.copy()) + ) + info = f"with {escape_single_quotes(parameters) or 'no parameters'} : {escape_single_quotes(content)}" + else: + self.warn_or_stop("Ignored negative command in negative prompt") + self.__visit(tree.children[1::]) + t2 = time.monotonic_ns() + self.__debug_end("commandstn", start_result, t2 - t1, info) + + def commandstni(self, tree: lark.Tree): + """ + Process a send to negative insertion point command in the tree and add it to the list of negative tags. + """ + start_result = self.result + info = None + t1 = time.monotonic_ns() + if self.__is_negative: + negtagparameters = tree.children[0] + if negtagparameters is not None: + parameters = str(negtagparameters) + else: + parameters = "" + self.__negtags.append(self.NegTag(len(self.result), len(self.result), "", parameters, self.__shell.copy())) + info = f"with {parameters or 'no parameters'}" + else: + self.warn_or_stop("Ignored negative insertion point command in positive prompt") + t2 = time.monotonic_ns() + self.__debug_end("commandstni", start_result, t2 - t1, info) + + def __varset( + self, + command: str, + variable: str, + modifiers: lark.Tree | None, + content: lark.Tree | None, + ): + """ + Process a generic set command in the tree. + """ + t1 = time.monotonic_ns() + start_result = self.result + if variable.startswith("_"): + self.warn_or_stop( + f"Invalid variable name '{escape_single_quotes(variable)}' detected! System variables cannot be set." + ) + return + info = variable + value_description = self.__get_original_node_content(content, None) + value = content + modifiers_str: list[str] = [str(m) for m in modifiers.children] if modifiers is not None else [] + if any(item in modifiers_str for item in ["+", "add"]): + info += f" += '{escape_single_quotes(value_description or '')}'" + raw_oldvalue = self.__state.user_variables.get(variable, None) + if raw_oldvalue is None: + newvalue = value + self.warn_or_stop(f"Unknown variable {escape_single_quotes(variable)}") + elif isinstance(raw_oldvalue, str): + newvalue = lark.Tree( + lark.Token("RULE", "varvalue"), + [lark.Token("plain", raw_oldvalue), value], + # Meta should be {"content": raw_oldvalue + value}, + ) + else: + newvalue = lark.Tree( + lark.Token("RULE", "varvalue"), + [raw_oldvalue, value], + # Meta should be {"content": raw_oldvalue.meta.content + value.meta.content}, + ) + elif any(item in modifiers_str for item in ["?", "ifundefined"]): + info += f" ?= '{escape_single_quotes(value_description or '')}'" + raw_oldvalue = self.__state.user_variables.get(variable, None) + if raw_oldvalue is None: + newvalue = value + else: + info += " (not set)" + newvalue = None + else: + newvalue = value + if newvalue is not None: + if any(item in modifiers_str for item in ["!", "evaluate"]): + newvalue = self.__visit(newvalue, False, True) + info += " =! " + else: + info += " = " + self.__set_user_variable_value(variable, newvalue) + currentvalue = self.__get_user_variable_value(variable, False) + if currentvalue is None: + info += "not evaluated yet" + else: + info += f"'{escape_single_quotes(currentvalue)}'" + t2 = time.monotonic_ns() + self.__debug_end(command, start_result, t2 - t1, info) + + def variableset(self, tree: lark.Tree): + """ + Process a DP set variable command in the tree and add it to the dictionary of variables. + """ + modifiers = tree.children[1] or lark.Tree(lark.Token("RULE", "variablesetmodifiers"), []) + immediate = tree.children[2] + if immediate is not None: + modifiers.children = modifiers.children.copy() + modifiers.children.append(immediate) + self.__varset("variableset", str(tree.children[0]), modifiers, tree.children[3]) + + def commandset(self, tree: lark.Tree): + """ + Process a set command in the tree and add it to the dictionary of variables. + """ + self.__varset("commandset", str(tree.children[0]), tree.children[1], tree.children[2]) + + def __varecho(self, command: str, variable: str, default: lark.Tree | None): + """ + Process a generic echo command in the tree. + """ + t1 = time.monotonic_ns() + start_result = self.result + default_value = None + # if default is not None: + # default_value = self.__visit(default, True) # for log + value = self.__get_user_variable_value(variable, True, True) + if value is None: + if default is not None: + if self.__debug_level == DEBUG_LEVEL.full: + self.__state.logger.debug( + f"Variable '{escape_single_quotes(variable)}' not found, using default value" + ) + v = self.__visit(default, False, True) + self.result += v + default_value = v + self.__state.echoed_variables[variable] = v + else: + self.warn_or_stop(f"Unknown variable {escape_single_quotes(variable)}") + default_value = "" + self.__state.echoed_variables[variable] = "" + else: + self.__state.echoed_variables[variable] = value + t2 = time.monotonic_ns() + info = variable + if default_value is not None: + info += f" with default '{escape_single_quotes(default_value)}'" + self.__debug_end(command, start_result, t2 - t1, info) + + def variableuse(self, tree: lark.Tree): + """ + Process a DP use variable command in the tree. + """ + self.__varecho("variableuse", str(tree.children[0]), tree.children[1] if len(tree.children) > 1 else None) + + def commandecho(self, tree: lark.Tree): + """ + Process an echo command in the tree. + """ + self.__varecho("commandecho", str(tree.children[0]), tree.children[1] if len(tree.children) > 1 else None) + + def commandif(self, tree: lark.Tree): + """ + Process an if command in the tree. + """ + t1 = time.monotonic_ns() + start_result = self.result + for i, n in enumerate(tree.children): + content = n.children[-1] + if len(n.children) == 2: # its not an else + # has a condition + condition = n.children[0] + c = self.__get_original_node_content(condition, f"condition {i}") + if self.__eval_condition(condition): + self.__visit(content) + t2 = time.monotonic_ns() + self.__debug_end("commandif", start_result, t2 - t1, c) + return + else: # its an else + self.__visit(content) + t2 = time.monotonic_ns() + self.__debug_end("commandif", start_result, t2 - t1, "else") + return + + def commandext(self, tree: lark.Tree): + """ + Process an extranetwork command in the tree. + """ + t1 = time.monotonic_ns() + start_result = self.result + extnet = "(ignored)" + if not self.__state.options.rem_removeextranetworktags: + extnet_type: str = (tree.children[0].children[0] or "") + str(tree.children[0].children[1]) + is_mapping = extnet_type.startswith("$") + if is_mapping: + extnet_type = extnet_type[1:] + extnet_id: str = str(tree.children[1]) + if extnet_id.startswith("'") or extnet_id.startswith('"'): + extnet_id = extnet_id[1:-1] + extnet_id = re.sub(r"\\(.)", r"\1", extnet_id) # so we can escape some special characters + parameters: str = "" + parameters_defaulted = False + if tree.children[2]: + parameters = str(tree.children[2]) + elif extnet_type in ("lora", "hypernet"): + parameters = "1" + parameters_defaulted = True + if parameters.startswith("'") or parameters.startswith('"'): + parameters = parameters[1:-1] + parameters_is_number = bool(re.match(r"^[-+]?\d*\.?\d+$", parameters or "")) + condition = tree.children[3] + if not condition or self.__eval_condition(condition): + extnet_id = f"{extnet_type}:{extnet_id}" + triggers = tree.children[4] if len(tree.children) > 4 else None + extra_triggers = None + compiled_extra_triggers = None + if is_mapping: + found = self.__state.extranetwork_mappings_obj.cached_mappings.get(extnet_id, None) + # we assume the conditions do not change inside the prompt + found_in_cache = found is not None + if found is None: + found_mappings: list[PPPENMappingVariant] = [] + else_mapping = None + if self.__state.extranetwork_mappings_obj: + enmapping = self.__state.extranetwork_mappings_obj.extranetwork_mappings.get( + extnet_id, None + ) + if enmapping: + for v in enmapping.variants: + if v.condition: + try: + cnd = parse_prompt( + self.__state, + "condition", + v.condition, + self.__state.parsers["condition"], + True, + ) + except lark.exceptions.UnexpectedInput as e: + self.warn_or_stop( + f"Error parsing condition '{escape_single_quotes(v.condition)}' in extranetwork mapping '{escape_single_quotes(extnet_id)}'! : {e.__class__.__name__}", + e, + ) + cnd = None + else: + cnd = "True" + if cnd is not None and (cnd == "True" or self.__eval_condition(cnd)): + if v.condition: + found_mappings.append(v) + else: + else_mapping = v + if found_mappings: + found = found_mappings[ + self.__rng.choice( + len(found_mappings), + p=[v.weight or 1 for v in found_mappings], + ) + ] + else: + found = else_mapping + self.__state.extranetwork_mappings_obj.cached_mappings[extnet_id] = found + if found: + if found.name: + if not found_in_cache and self.__debug_level != DEBUG_LEVEL.none: + self.__state.logger.info( + f"Mapping extranetwork '{escape_single_quotes(extnet_id)}' to '{escape_single_quotes(extnet_type)}:{escape_single_quotes(found.name)}'" + ) + extnet_id = f"{extnet_type}:{found.name}" + f_parameters = found.parameters + if not f_parameters and extnet_type in ("lora", "hypernet"): + f_parameters = "1" + found_parameters_is_number = True + else: + found_parameters_is_number = f_parameters and bool( + re.match(r"^[-+]?\d*\.?\d+$", str(f_parameters) or "") + ) + if found_parameters_is_number and parameters_is_number: + parameters = f"{float(f_parameters) * float(parameters):.2f}".rstrip("0").rstrip(".") + elif f_parameters is not None and parameters_defaulted: + parameters = f_parameters + elif found.triggers: + if not found_in_cache and self.__debug_level != DEBUG_LEVEL.none: + self.__state.logger.info( + f"Mapping extranetwork '{escape_single_quotes(extnet_id)}' to just triggers" + ) + extnet_id = None + else: + if not found_in_cache and self.__debug_level != DEBUG_LEVEL.none: + self.__state.logger.info( + f"Mapping extranetwork '{escape_single_quotes(extnet_id)}' to nothing" + ) + extnet_id = None + if found.triggers: + extra_triggers = ", ".join(found.triggers) + try: + compiled_extra_triggers = parse_prompt( + self.__state, "triggers", extra_triggers, self.__state.parsers["content"], True + ) + except lark.exceptions.UnexpectedInput as e: + self.warn_or_stop( + f"Error parsing triggers '{escape_single_quotes(extra_triggers)}' in extranetwork mapping '{escape_single_quotes(extnet_id)}'! : {e.__class__.__name__}", + e, + ) + compiled_extra_triggers = None + else: + self.warn_or_stop(f"Extranetwork mapping '{escape_single_quotes(extnet_id)}' not found!") + if extnet_id: + extnet = f"<{extnet_id}:{parameters}>" + self.result += extnet + elif triggers or compiled_extra_triggers: + extnet = "(only triggers)" + if triggers or compiled_extra_triggers: + if extnet_id: + if not self.__state.options.cup_extranetworktags: + self.result += " " + else: + self.result += ", " + if triggers: + self.result += self.__visit(triggers, True, True) + if compiled_extra_triggers: + if triggers: + self.result += ", " + self.result += self.__visit(compiled_extra_triggers, True, True) + if triggers or compiled_extra_triggers: + self.result += ", " + t2 = time.monotonic_ns() + self.__debug_end("commandext", start_result, t2 - t1, extnet) + + def commandsetwcdeffilter(self, tree: lark.Tree): + """ + Process a setwcdeffilter (Set Wildcard Default Filter) command in the tree. + """ + t1 = time.monotonic_ns() + start_result = self.result + wildcard_key: str = self.__visit(tree.children[0].children[1], False, True) + selected_wildcards = [x.key for x in self.__state.wildcards_obj.get_wildcards(wildcard_key)] + if not selected_wildcards: + self.warn_or_stop(f"Wildcard '{escape_single_quotes(wildcard_key)}' not found for default filter setting!") + else: + filter_object = tree.children[1].children[1] if tree.children[1] is not None else None + if filter_object is None: + for wc in selected_wildcards: + if self.__debug_level == DEBUG_LEVEL.full: + self.__state.logger.debug(f"Removed default filter for wildcard '{escape_single_quotes(wc)}'") + self.__state.wildcards_obj.set_wildcard_default_filter(wc, None) + else: + filter_specifier = self.__extract_filter_specifiers(filter_object) + for wc in selected_wildcards: + if self.__debug_level == DEBUG_LEVEL.full: + self.__state.logger.debug(f"Set default filter for wildcard '{escape_single_quotes(wc)}'") + self.__state.wildcards_obj.set_wildcard_default_filter(wc, filter_specifier) + t2 = time.monotonic_ns() + self.__debug_end("commandsetwcdeffilter", start_result, t2 - t1) + + def extranetworktag(self, tree: lark.Tree): + """ + Process an extra network construct in the tree. + """ + t1 = time.monotonic_ns() + start_result = self.result + if not self.__state.options.rem_removeextranetworktags: + self.result += f"<{tree.children[0]}" + self.__visit(tree.children[1]) + self.result += ">" + t2 = time.monotonic_ns() + self.__debug_end("extranetworktag", start_result, t2 - t1) + + def __get_choices_internal_get( + self, + choice_values: list[dict], + filter_specifier: Optional[list[list[str]]] = None, + wildcard_key: str = None, + ) -> list[dict]: + msg_where = f"wildcard '{escape_single_quotes(wildcard_key)}'" if wildcard_key else "choices" + if filter_specifier is not None: + filtered_choice_values = [] + for i, c in enumerate(choice_values): + passes = False + for or_ in filter_specifier: + tmp_pass = True + for and_ in or_: + if and_.isdecimal(): + if int(and_) != i: + tmp_pass = False + break + elif re.match(r"^\d+-\d+$", and_): + try: + start, end = map(int, and_.split("-", 1)) + if not start <= i <= end: + tmp_pass = False + break + except ValueError: + tmp_pass = False + break + elif and_.lower() not in c.get("labels", []): + tmp_pass = False + break + if tmp_pass: + passes = True + break + if passes: + filtered_choice_values.append(c) + if not filtered_choice_values: + self.warn_or_stop( + f"Wildcard filter specifier '{escape_single_quotes(','.join(['+'.join(y for y in x) for x in filter_specifier]))}' found no matches in choices for wildcard '{escape_single_quotes(wildcard_key)}'!" + ) + else: + filtered_choice_values = choice_values.copy() + expanded_choice_values = [] + for i, c in enumerate(filtered_choice_values): + if c.get("command", False): + content_text = self.__visit(c.get("content", ""), False, True).strip() + (cmd, cmd_args) = content_text.split() + if cmd == "include": + wcs = self.__state.wildcards_obj.get_wildcards(cmd_args) + if not wcs: + self.warn_or_stop( + f"Not found included wildcard '{escape_single_quotes(cmd_args)}' at {msg_where}!" + ) + c_weight = float(c.get("weight", 1.0)) + for wc in wcs: + if wc.key in self.__seen_wildcards: + self.warn_or_stop( + f"Circular reference detected including wildcard '{escape_single_quotes(wc.key)}' at {msg_where} (chain starts at '{escape_single_quotes(self.__seen_wildcards[0])}')!" + ) + continue + self.__seen_wildcards.append(wc.key) + if self.__debug_level == DEBUG_LEVEL.full: + self.__state.logger.debug(f"Seen wildcard '{escape_single_quotes(wc.key)}'") + self.__state.logger.debug( + f"Including choices from wildcard '{escape_single_quotes(wc.key)}'" + ) + (_, choice_values) = self.__check_wildcard_initialization(wc) + if choice_values is not None: + ch_values = self.__get_choices_internal_get(choice_values, None, wc.key) + for cv in ch_values: + expanded_choice_values.append( + { + **cv, + "weight": float(cv.get("weight", 1.0) * c_weight), # we adjust the weight + } + ) + else: + self.warn_or_stop(f"Unsupported choice command '{escape_single_quotes(cmd)}' at {msg_where}!") + else: + expanded_choice_values.append(c) + return expanded_choice_values + + def __get_choices_internal_select( + self, + options: dict | None, + choice_values: list[dict], + filter_specifier: Optional[list[list[str]]] = None, + wildcard_key: str = None, + ) -> tuple[str, list[str], str, str]: + """ + Select choices based on the options. + + Args: + options (dict): The object representing the options construct. + choice_values (list[dict]): A list of choice objects. + filter_specifier (list[list[str]]): The filter specifier. + wildcard_key (str): The wildcard key if it is a wildcard. + + Returns: + tuple: A tuple containing the prefix, selected choices, separator and suffix + """ + seen_wildcards_len = len(self.__seen_wildcards) + if options is None: + options = {} + sampler: str = options.get("sampler", "~") + repeating: bool = options.get("repeating", False) + optional: bool = options.get("optional", False) + if "count" in options: + from_value = options["count"] + to_value = from_value + else: + from_value: int = options.get("from", 1) + to_value: int = options.get("to", 1) + separator: str = options.get("separator", self.__state.options.wil_choice_separator) + msg_where = f"wildcard '{escape_single_quotes(wildcard_key)}'" if wildcard_key else "choices" + if sampler != "~": + self.warn_or_stop(f"Unsupported sampler '{escape_single_quotes(sampler)}' at {msg_where} options!") + sampler = "~" + expanded_choice_values = self.__get_choices_internal_get(choice_values, filter_specifier, wildcard_key) + available_choices: list[dict] = [] + weights = [] + included_choices = 0 + excluded_choices = 0 + excluded_weights_sum = 0 + for i, c in enumerate(expanded_choice_values): + c["choice_index"] = i # we index them to later sort the results + weight = float(c.get("weight", 1.0)) + condition = c.get("if", None) + if weight > 0 and (condition is None or self.__eval_condition(condition)): + available_choices.append(c) + weights.append(weight) + included_choices += 1 + else: + weights.append(-1) + excluded_choices += 1 + excluded_weights_sum += weight + if excluded_choices > 0: # we need to redistribute the excluded weights + weights = [weight + excluded_weights_sum / included_choices for weight in weights if weight >= 0] + weights = np.array(weights) + weights /= weights.sum() # normalize weights + if available_choices: + if from_value < 0: + from_value = 1 + elif from_value > len(available_choices): + from_value = len(available_choices) + if to_value < 1: + to_value = 1 + elif (to_value > len(available_choices) and not repeating) or from_value > to_value: + to_value = len(available_choices) + num_choices = ( + self.__rng.integers(from_value, to_value, endpoint=True) if from_value < to_value else from_value + ) + else: + num_choices = 0 + if not optional and from_value > 0: + self.warn_or_stop(f"Not enough choices found for {msg_where}!") + if num_choices < 2: + repeating = False + if self.__debug_level == DEBUG_LEVEL.full: + self.__state.logger.debug( + format_output( + f"Selecting {'optional ' if optional else ''}{'repeating ' if repeating else ''}{num_choices} choice" + + ("s" if num_choices != 1 else "") + + (f" and separating with '{escape_single_quotes(separator)}'" if num_choices > 1 else "") + ) + ) + if num_choices > 0: + selected_choices: list[dict] = ( + list(self.__rng.choice(available_choices, size=num_choices, p=weights, replace=repeating)) + if available_choices + else [] + ) + if self.__state.options.wil_keep_choices_order: + selected_choices = sorted(selected_choices, key=lambda x: x["choice_index"]) + selected_choices_text = [] + prefix: str = ( + self.__visit(options.get("prefix", None), False, True) + if options.get("prefix", None) is not None + else "" + ) + if prefix != "" and re.match(r"\w", prefix[-1]): + prefix += " " + for i, c in enumerate(selected_choices): + t1 = time.monotonic_ns() + choice_content_obj = c.get("content", c.get("text", None)) + if isinstance(choice_content_obj, str): + choice_content = choice_content_obj + else: + choice_content = self.__visit(choice_content_obj, False, True) + t2 = time.monotonic_ns() + if self.__debug_level == DEBUG_LEVEL.full: + self.__state.logger.debug( + f"Adding choice {i+1} ({(t2-t1) / 1_000_000_000:.3f} seconds):\n" + + textwrap.indent(re.sub(r"\n$", "", choice_content), " ") + ) + selected_choices_text.append(choice_content) + suffix: str = ( + self.__visit(options.get("suffix", None), False, True) + if options.get("suffix", None) is not None + else "" + ) + if suffix != "" and re.match(r"\w", suffix[0]): + suffix = " " + suffix + # remove comments + results = [re.sub(r"\s*#[^\n]*(?:\n|$)", "", r, flags=re.DOTALL) for r in selected_choices_text] + else: + prefix = "" + suffix = "" + results = [] + if self.__debug_level == DEBUG_LEVEL.full: + list_unseen = [f"'{escape_single_quotes(x)}'" for x in self.__seen_wildcards[seen_wildcards_len:]] + self.__state.logger.debug(f"Unseen wildcards: {', '.join(list_unseen)}") + self.__seen_wildcards = self.__seen_wildcards[:seen_wildcards_len] + return (prefix, results, separator, suffix) + + def __get_choices( + self, + options: dict | None, + choice_values: list[dict], + filter_specifier: Optional[list[list[str]]] = None, + wildcard_key: str = None, + ) -> str: + r = self.__get_choices_internal_select(options, choice_values, filter_specifier, wildcard_key) + if r[1]: + return r[0] + r[2].join(r[1]) + r[3] + return "" + + def __convert_choices_options(self, options: Optional[lark.Tree], is_wcdef: bool = False) -> dict: + """ + Convert the choices options to a dictionary. + + Args: + options (Tree): The choices options tree. + + Returns: + dict: The converted choices options. + """ + if options is None: + return None + options_dict = {} + if len(options.children) == 1: + if options.children[0] is not None: + options_dict["sampler"] = str(options.children[0]) + else: + if options.children[0] is not None: + options_dict["sampler"] = str(options.children[0].children[0]) + if options.children[1] is not None: + flags = str(options.children[1].children[0]) + options_dict["repeating"] = "r" in flags + options_dict["optional"] = "o" in flags + irange = 2 + idesc = 3 + isep = 4 + if options.children[irange] is not None: + if len(options.children[irange].children) == 1: + if options.children[irange].children[0] is not None: + options_dict["count"] = int(options.children[irange].children[0]) + else: + options_dict["from"] = ( + int(options.children[irange].children[0]) + if options.children[irange].children[0] is not None + else 1 + ) + options_dict["to"] = ( + int(options.children[irange].children[1]) + if options.children[irange].children[1] is not None + else 1 + ) + if is_wcdef: + if options.children[idesc] is not None: + options_dict["description"] = str(options.children[idesc].children[0])[1:-1] + else: + isep -= 1 # only wildcard definition options have a description + if options.children[isep] is not None: + options_dict["separator"] = self.__visit(options.children[isep], False, True) + if not options_dict: + options_dict = None + return options_dict + + def __convert_choice(self, choice: lark.Tree) -> dict: + """ + Convert the choice to a dictionary. + + Args: + choice (Tree): The choice tree. + + Returns: + dict: The converted choice. + """ + choice_dict = {} + choice_dict["command"] = choice.children[0] is not None + c_label_obj = choice.children[1] + choice_dict["labels"] = [str(x).lower() for x in c_label_obj.children[1:-1]] if c_label_obj is not None else [] + choice_dict["weight"] = float(choice.children[2].children[0]) if choice.children[2] is not None else 1.0 + choice_dict["if"] = choice.children[3].children[0] if choice.children[3] is not None else None + choice_dict["content"] = choice.children[-1] + return choice_dict + + def __check_wildcard_initialization(self, wildcard: PPPWildcard) -> tuple[dict | None, list[dict] | None]: + """ + Initializes a wildcard if it hasn't been yet. + + Args: + wildcard (PPPWildcard): The wildcard to check. + Returns: + tuple: A tuple containing the options and choice values of the wildcard. + """ + choice_values = wildcard.choices + options = wildcard.options + if choice_values is None: + t1 = time.monotonic_ns() + choice_values = [] + options, n = self.get_wildcard_options(wildcard) + # we process the choices + for cv in wildcard.unprocessed_choices[n:]: + if isinstance(cv, dict): + if self.__state.wildcards_obj.is_dict_choice_options(cv): + condition = cv.get("if", None) + if condition is not None and isinstance(condition, str): + try: + cv["if"] = parse_prompt( + self.__state, "condition", condition, self.__state.parsers["condition"], True + ) + except lark.exceptions.UnexpectedInput as e: + self.warn_or_stop( + f"Error parsing condition '{escape_single_quotes(condition)}' in wildcard '{escape_single_quotes(wildcard.key)}'! : {e.__class__.__name__}", + e, + ) + cv["if"] = None + content = cv.get("content", cv.get("text", None)) + cv["content"] = content + if "text" in cv: + del cv["text"] + if content is not None and isinstance(content, str): + try: + cv["content"] = parse_prompt( + self.__state, "choicevalue", content, self.__state.parsers["choicevalue"], True + ) + except lark.exceptions.UnexpectedInput as e: + self.warn_or_stop( + f"Error parsing choice content '{escape_single_quotes(content)}' in wildcard '{escape_single_quotes(wildcard.key)}'! : {e.__class__.__name__}", + e, + ) + cv["content"] = None + if cv["content"] is not None: + if self.__debug_level == DEBUG_LEVEL.full: + self.__state.logger.debug(f"Processed choice {cv}") + choice_values.append(cv) + else: + self.warn_or_stop( + f"Invalid choice {cv} in wildcard '{escape_single_quotes(wildcard.key)}'!" + ) + else: + self.warn_or_stop(f"Invalid choice {cv} in wildcard '{escape_single_quotes(wildcard.key)}'!") + else: + try: + choice_values.append( + self.__convert_choice( + parse_prompt(self.__state, "choice", cv, self.__state.parsers["choice"], True) + ) + ) + except lark.exceptions.UnexpectedInput as e: + self.warn_or_stop( + f"Error parsing choice '{escape_single_quotes(cv)}' in wildcard '{escape_single_quotes(wildcard.key)}'! : {e.__class__.__name__}", + e, + ) + wildcard.choices = choice_values + t2 = time.monotonic_ns() + if self.__debug_level == DEBUG_LEVEL.full: + self.__state.logger.debug( + f"Processed choices for wildcard '{escape_single_quotes(wildcard.key)}' ({(t2-t1) / 1_000_000_000:.3f} seconds)" + ) + return (self.__clean_wildcard_options(options), choice_values) + + def get_wildcard_options(self, wildcard: PPPWildcard) -> tuple[dict | None, int]: + options = wildcard.options + n = 0 + # we check the first choice to see if it is actually options + if isinstance(wildcard.unprocessed_choices[0], dict): + if self.__state.wildcards_obj.is_dict_wcdef_options(wildcard.unprocessed_choices[0]): + options = wildcard.unprocessed_choices[0] + prefix = options.get("prefix", None) + if prefix is not None and isinstance(prefix, str): + try: + options["prefix"] = parse_prompt( + self.__state, "choicevalue", prefix, self.__state.parsers["choicevalue"], True + ) + except lark.exceptions.UnexpectedInput as e: + self.warn_or_stop( + f"Error parsing choice prefix '{escape_single_quotes(prefix)}' in wildcard '{escape_single_quotes(wildcard.key)}'! : {e.__class__.__name__}", + e, + ) + suffix = options.get("suffix", None) + if suffix is not None and isinstance(suffix, str): + try: + options["suffix"] = parse_prompt( + self.__state, "choicevalue", suffix, self.__state.parsers["choicevalue"], True + ) + except lark.exceptions.UnexpectedInput as e: + self.warn_or_stop( + f"Error parsing choice suffix '{escape_single_quotes(suffix)}' in wildcard '{escape_single_quotes(wildcard.key)}'! : {e.__class__.__name__}", + e, + ) + n = 1 + else: + if wildcard.unprocessed_choices[0].endswith("$$"): + try: + options = self.__convert_choices_options( + parse_prompt( + self.__state, + "as wildcard options", + wildcard.unprocessed_choices[0][:-2].strip(), + self.__state.parsers["wcdefoptions"], + True, + ), + True, + ) + n = 1 + except lark.exceptions.UnexpectedInput: + options = None + if options is None and self.__debug_level == DEBUG_LEVEL.full: + self.__state.logger.debug("Does not have options") + wildcard.options = options + return options, n + + def __clean_wildcard_options(self, options: dict) -> dict: + if options is not None: + options.pop("description", None) # description is only for wildcard definitions, not for usage + if not options: + options = None + return options + + def wildcard(self, tree: lark.Tree): + """ + Process a wildcard construct in the tree. + """ + t1 = time.monotonic_ns() + seen_wildcards_len = len(self.__seen_wildcards) + start_result = self.result + applied_options = self.__clean_wildcard_options(self.__convert_choices_options(tree.children[0], False)) + wildcard_key: str = self.__visit(tree.children[1], False, True) + wc = self.__get_original_node_content(tree, f"?__{wildcard_key}__") + if self.__state.options.wil_process_wildcards: + if self.__debug_level == DEBUG_LEVEL.full: + self.__state.logger.debug(f"Processing wildcard: {wildcard_key}") + selected_wildcards = self.__state.wildcards_obj.get_wildcards(wildcard_key) + if not selected_wildcards: + self.detectedWildcards.append(wc) + self.result += wc + t2 = time.monotonic_ns() + self.__debug_end("wildcard", start_result, t2 - t1, wc) + return + filter_specifier: list[list[str]] = None + filter_object = tree.children[2] + if filter_object is not None: + if ( # it's an inherited filter from another wildcard + isinstance(filter_object.children[1], lark.Token) + and filter_object.children[1] is not None + and "^" in str(filter_object.children[1]) + ): + filter_wildcard_key = self.__visit(filter_object.children[2], False, True) + filter_specifier = self.__wildcard_filters.get(filter_wildcard_key, None) + if self.__debug_level == DEBUG_LEVEL.full: + self.__state.logger.debug("Filtering choices with inherited filter") + else: + filter_specifier = self.__extract_filter_specifiers(filter_object.children[2]) + if self.__debug_level == DEBUG_LEVEL.full: + self.__state.logger.debug("Filtering choices") + self.__wildcard_filters[wildcard_key] = filter_specifier + if filter_object.children[1] is not None and "#" in str( + filter_object.children[1] + ): # means do not use the filter in this wildcard + if self.__debug_level == DEBUG_LEVEL.full: + self.__state.logger.debug("Ignoring filter") + filter_specifier = None + else: + filter_specifier = self.__state.wildcards_obj.get_wildcard_default_filter(wildcard_key) + if filter_specifier is not None: + self.__wildcard_filters[wildcard_key] = filter_specifier + if self.__debug_level == DEBUG_LEVEL.full: + self.__state.logger.debug("Applying default filter") + if ( + len(selected_wildcards) > 1 + and filter_specifier is not None + and any(y[0].isdecimal() for x in filter_specifier for y in x) + ): + self.__state.logger.warning( + f"Using a globbing wildcard '{escape_single_quotes(wildcard_key)}' with positional index filters is not recommended!" + ) + var_object = tree.children[3] + variablename = None + variablebackup = None + if var_object is not None: + variablename = str(var_object.children[0]) + 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) + choice_values_all = [] + for wildcard in selected_wildcards: + if wildcard is None: + self.detectedWildcards.append(wc) + self.result += wc + t2 = time.monotonic_ns() + self.__debug_end("wildcard", start_result, t2 - t1, wc) + return + if wildcard.key in self.__seen_wildcards: + self.warn_or_stop( + f"Circular reference detected with wildcard '{escape_single_quotes(self.__seen_wildcards[-1])}' (chain starts at '{escape_single_quotes(self.__seen_wildcards[0])}')!" + ) + continue + self.__seen_wildcards.append(wildcard.key) + if self.__debug_level == DEBUG_LEVEL.full: + self.__state.logger.debug(f"Seen wildcard '{escape_single_quotes(wildcard.key)}'") + (options, choice_values) = self.__check_wildcard_initialization(wildcard) + if options is not None: + if applied_options is None: + applied_options = options + else: + if self.__debug_level == DEBUG_LEVEL.full: + self.__state.logger.debug( + f"Options for wildcard '{escape_single_quotes(wildcard.key)}' are ignored!" + ) + choice_values_all += choice_values + self.result += self.__get_choices(applied_options, choice_values_all, filter_specifier, wildcard_key) + if wildcard_key in self.__wildcard_filters: + del self.__wildcard_filters[wildcard_key] + if variablename is not None: + self.__remove_user_variable(variablename) + if variablebackup is not None: + self.__state.user_variables[variablename] = variablebackup + elif self.__state.options.wil_ifwildcards != IFWILDCARDS_CHOICES.remove: + self.detectedWildcards.append(wc) + self.result += wc + if self.__debug_level == DEBUG_LEVEL.full: + list_unseen = [f"'{escape_single_quotes(x)}'" for x in self.__seen_wildcards[seen_wildcards_len:]] + self.__state.logger.debug(f"Unseen wildcards: {', '.join(list_unseen)}") + self.__seen_wildcards = self.__seen_wildcards[:seen_wildcards_len] + t2 = time.monotonic_ns() + self.__debug_end("wildcard", start_result, t2 - t1, f"'{escape_single_quotes(wc)}'") + + def __extract_filter_specifiers(self, filters: lark.Tree) -> list[list[str]]: + filter_specifier = [] + for or_ in filters.children: + for and_ in or_.children: + label = and_.children[0] + if isinstance(label, lark.Token): + # it's a literal, we can use it directly + filter_specifier.append([str(label)]) + else: + # it's a variable, we need to evaluate it + v = self.__visit(label, False, True) + filter_specifier.append([v]) + return filter_specifier + + def choices(self, tree: lark.Tree): + """ + Process a choices construct in the tree. + """ + t1 = time.monotonic_ns() + start_result = self.result + options = self.__convert_choices_options(tree.children[0], False) + choice_values = [self.__convert_choice(c) for c in tree.children[1::]] + ch = self.__get_original_node_content(tree, "?{...}") + if self.__state.options.wil_process_wildcards: + if self.__debug_level == DEBUG_LEVEL.full: + self.__state.logger.debug("Processing choices:") + self.result += self.__get_choices(options, choice_values) + elif self.__state.options.wil_ifwildcards != IFWILDCARDS_CHOICES.remove: + self.detectedWildcards.append(ch) + self.result += ch + t2 = time.monotonic_ns() + self.__debug_end("choices", start_result, t2 - t1, f"'{escape_single_quotes(ch)}'") + + def __default__(self, tree): + t1 = time.monotonic_ns() + start_result = self.result + self.__visit(tree.children) + t2 = time.monotonic_ns() + self.__debug_end(tree.data.value, start_result, t2 - t1) + + def start(self, tree): + self.result = "" + t1 = time.monotonic_ns() + self.__visit(tree.children) + attention_processing = self.__state.host_config.get("attention", "ok") + # process the found negative tags + for negtag in self.__negtags: + if self.__state.options.cup_mergeattention: + # join consecutive attention elements + for i in range(len(negtag.shell) - 1, 0, -1): + if negtag.shell[i].type == "at" and negtag.shell[i - 1].type == "at": + new_weight = ( # we limit the new weight to two decimals + math.floor(100 * float(negtag.shell[i - 1].data[1]) * float(negtag.shell[i].data[1])) / 100 + ) + new_weight_str = f"{new_weight:.2f}".rstrip("0").rstrip(".") + if new_weight_str == "0.9" and attention_processing != "parentheses": + new_kind = 1 + elif new_weight_str == "1.1": + new_kind = 2 + else: + new_kind = 3 + negtag.shell[i - 1] = self.AccumulatedShell( + "at", + (new_kind, new_weight_str), + ) + negtag.shell.pop(i) + start = "" + end = "" + for s in negtag.shell: + match s.type: + case "at": + if s.data[0] == 1: + start += "[" + end = "]" + end + elif s.data[0] == 2: + start += "(" + end = ")" + end + else: # 3 + start += "(" + end = f":{s.data[1]})" + end + # case "sc": + case "scb": + start += "[" + end = f"::{s.data}]" + end + case "sca": + start += "[" + end = f":{s.data}]" + end + # case "al": + case "alo": + start += "[" + ("|" * int(s.data["pos"] - 1)) + end = ("|" * int(s.data["len"] - s.data["pos"])) + "]" + end + content = start + negtag.content + end + position = negtag.parameters or "s" + if position.startswith("i"): + n = int(position[1]) + self.insertion_at[n] = [negtag.start, negtag.end] + elif len(content) > 0: + if content not in self.__already_processed: + if self.__state.options.stn_ignore_repeats: + self.__already_processed.append(content) + if self.__debug_level == DEBUG_LEVEL.full: + self.__state.logger.debug(format_output(f"Adding content at position {position}: {content}")) + if position == "e": + self.add_at["end"].append(content) + elif position.startswith("p"): + n = int(position[1]) + self.add_at["insertion_point"][n].append(content) + else: # position == "s" or invalid + self.add_at["start"].append(content) + else: + self.__state.logger.warning(format_output(f"Ignoring repeated content: {content}")) + t2 = time.monotonic_ns() + self.__debug_end("start", "", t2 - t1) diff --git a/ppp_utils.py b/ppp_utils.py index 18ef055..c032eb0 100644 --- a/ppp_utils.py +++ b/ppp_utils.py @@ -16,6 +16,7 @@ def deep_freeze(obj): return tuple(deep_freeze(i) for i in sorted(obj)) return obj + def escape_single_quotes(s: str): """ Escape single quotes in a string. @@ -28,6 +29,7 @@ def escape_single_quotes(s: str): """ return s.replace("'", "\\'") + def escape_double_quotes(s: str): """ Escape double quotes in a string. @@ -39,3 +41,16 @@ def escape_double_quotes(s: str): str: The escaped string. """ return s.replace('"', '\\"') + + +def format_output(text: str) -> str: + """ + Formats the output text by encoding it using unicode_escape and decoding it using utf-8. + + Args: + text (str): The input text to be formatted. + + Returns: + str: The formatted output text. + """ + return text.encode("unicode_escape").decode("utf-8") diff --git a/ppp_wildcards.py b/ppp_wildcards.py index ecc7ecc..08dbc75 100644 --- a/ppp_wildcards.py +++ b/ppp_wildcards.py @@ -52,7 +52,7 @@ class PPPWildcards: DEFAULT_WILDCARDS_FOLDER = "wildcards" LOCALINPUT_FILENAME = R"//INPUT\\" - def __init__(self, logger): + def __init__(self, logger=None): self.__logger: logging.Logger = logger self.__debug_level = DEBUG_LEVEL.none self.__wildcards_folders = [] @@ -77,7 +77,7 @@ class PPPWildcards: """ self.__debug_level = debug_level self.__wildcards_folders = wildcards_folders or [] - # if self.__debug_level != DEBUG_LEVEL.none: + # if self.__debug_level != DEBUG_LEVEL.none and self.__logger: # self.__logger.info("Refreshing wildcards...") # t1 = time.monotonic_ns() for fullpath in list(self.__wildcard_files.keys()): @@ -108,7 +108,7 @@ class PPPWildcards: self.wildcards = {} self.__wildcard_files = {} # t2 = time.monotonic_ns() - # if self.__debug_level != DEBUG_LEVEL.none: + # if self.__debug_level != DEBUG_LEVEL.none and self.__logger: # self.__logger.info(f"Wildcards refresh time: {(t2 - t1) / 1_000_000_000:.3f} seconds") def get_wildcards(self, key: str) -> list[PPPWildcard]: @@ -171,7 +171,7 @@ class PPPWildcards: debug (bool): Whether to print debug messages or not. """ last_modified_cached = self.__wildcard_files.get(full_path, None) # a time or a hash - if debug and last_modified_cached is not None and self.__debug_level != DEBUG_LEVEL.none: + if debug and last_modified_cached is not None and self.__debug_level != DEBUG_LEVEL.none and self.__logger: if full_path == self.LOCALINPUT_FILENAME: self.__logger.debug("Removing from memory wildcards from input") else: @@ -200,7 +200,7 @@ class PPPWildcards: if extension not in (".txt", ".json", ".yaml", ".yml"): return self.__remove_wildcards_from_path(full_path, False) - if last_modified_cached is not None and self.__debug_level != DEBUG_LEVEL.none: + if last_modified_cached is not None and self.__debug_level != DEBUG_LEVEL.none and self.__logger: self.__logger.debug(f"Updating wildcards from file: {full_path}") if extension == ".txt": self.__get_wildcards_in_text_file(full_path, base) @@ -208,7 +208,8 @@ class PPPWildcards: self.__get_wildcards_in_structured_file(full_path, base) self.__wildcard_files[full_path] = last_modified except Exception as e: # pylint: disable=broad-except - self.__logger.error(f"Error reading wildcard file '{escape_single_quotes(full_path)}': {e}") + if self.__logger: + self.__logger.error(f"Error reading wildcard file '{escape_single_quotes(full_path)}': {e}") def __get_wildcards_in_input(self, wildcards_input: str): """ @@ -223,20 +224,22 @@ class PPPWildcards: if h == new_h: return self.__remove_wildcards_from_path(self.LOCALINPUT_FILENAME, False) - if h is not None and self.__debug_level != DEBUG_LEVEL.none: + if h is not None and self.__debug_level != DEBUG_LEVEL.none and self.__logger: self.__logger.debug("Updating wildcards from input") wildcards_input = wildcards_input.strip() if wildcards_input != "": try: content = yaml.safe_load(wildcards_input) except yaml.YAMLError as e: - self.__logger.warning(f"Invalid format for input wildcards: {e}") + if self.__logger: + self.__logger.warning(f"Invalid format for input wildcards: {e}") return if content is not None: self.__add_wildcard(content, self.LOCALINPUT_FILENAME, [self.LOCALINPUT_FILENAME]) self.__wildcard_files[self.LOCALINPUT_FILENAME] = new_h except Exception as e: # pylint: disable=broad-except - self.__logger.error(f"Error reading wildcards input: {e}") + if self.__logger: + self.__logger.error(f"Error reading wildcards input: {e}") # NOTE wcdef and choice options should not have properties in common @@ -251,7 +254,19 @@ class PPPWildcards: bool: Whether the dictionary is a valid wildcard definition options dictionary or not. """ return all( - k in ["sampler", "repeating", "optional", "count", "from", "to", "prefix", "suffix", "description", "separator"] + k + in [ + "sampler", + "repeating", + "optional", + "count", + "from", + "to", + "prefix", + "suffix", + "description", + "separator", + ] for k in d.keys() ) @@ -334,9 +349,10 @@ class PPPWildcards: # we assume it is an anonymous wildcard with options firstkey = list(c.keys())[0] return self.__create_anonymous_wildcard(full_path, key_parts, i, c[firstkey], firstkey) - self.__logger.warning( - f"Invalid choice {i+1} in wildcard '{escape_single_quotes('/'.join(key_parts))}' in file '{escape_single_quotes(full_path)}'!" - ) + if self.__logger: + self.__logger.warning( + f"Invalid choice {i+1} in wildcard '{escape_single_quotes('/'.join(key_parts))}' in file '{escape_single_quotes(full_path)}'!" + ) return None def __create_anonymous_wildcard(self, full_path, key_parts, i, content, options=None): @@ -378,20 +394,23 @@ class PPPWildcards: tmp_key_parts.extend(key.split("/")) fullkey = "/".join(tmp_key_parts) if self.wildcards.get(fullkey, None) is not None: - self.__logger.warning( - f"Duplicate wildcard '{escape_single_quotes(fullkey)}' in file '{escape_single_quotes(full_path)}' and '{escape_single_quotes(self.wildcards[fullkey].file)}'!" - ) + if self.__logger: + self.__logger.warning( + f"Duplicate wildcard '{escape_single_quotes(fullkey)}' in file '{escape_single_quotes(full_path)}' and '{escape_single_quotes(self.wildcards[fullkey].file)}'!" + ) else: obj = self.__get_nested(content, key) choices = self.__get_choices(obj, full_path, tmp_key_parts) if choices is None: - self.__logger.warning( - f"Invalid wildcard '{escape_single_quotes(fullkey)}' in file '{escape_single_quotes(full_path)}'!" - ) + if self.__logger: + self.__logger.warning( + f"Invalid wildcard '{escape_single_quotes(fullkey)}' in file '{escape_single_quotes(full_path)}'!" + ) elif fullkey.startswith("_"): - self.__logger.warning( - f"Invalid wildcard name '{escape_single_quotes(fullkey)}' in file '{escape_single_quotes(full_path)}'! (cannot start with underscore)" - ) + if self.__logger: + self.__logger.warning( + f"Invalid wildcard name '{escape_single_quotes(fullkey)}' in file '{escape_single_quotes(full_path)}'! (cannot start with underscore)" + ) else: self.wildcards[fullkey] = PPPWildcard(full_path, fullkey, choices) return @@ -400,23 +419,27 @@ class PPPWildcards: elif isinstance(content, (int, float, bool)): content = [str(content)] if not isinstance(content, list): - self.__logger.warning(f"Invalid wildcard in file '{escape_single_quotes(full_path)}'!") + if self.__logger: + self.__logger.warning(f"Invalid wildcard in file '{escape_single_quotes(full_path)}'!") return fullkey = "/".join(key_parts) if self.wildcards.get(fullkey, None) is not None: - self.__logger.warning( - f"Duplicate wildcard '{escape_single_quotes(fullkey)}' in file '{escape_single_quotes(full_path)}' and '{escape_single_quotes(self.wildcards[fullkey].file)}'!" - ) + if self.__logger: + self.__logger.warning( + f"Duplicate wildcard '{escape_single_quotes(fullkey)}' in file '{escape_single_quotes(full_path)}' and '{escape_single_quotes(self.wildcards[fullkey].file)}'!" + ) else: choices = self.__get_choices(content, full_path, key_parts) if choices is None: - self.__logger.warning( - f"Invalid wildcard '{escape_single_quotes(fullkey)}' in file '{escape_single_quotes(full_path)}'!" - ) + if self.__logger: + self.__logger.warning( + f"Invalid wildcard '{escape_single_quotes(fullkey)}' in file '{escape_single_quotes(full_path)}'!" + ) elif fullkey.startswith("_"): - self.__logger.warning( - f"Invalid wildcard name '{escape_single_quotes(fullkey)}' in file '{escape_single_quotes(full_path)}'! (cannot start with underscore)" - ) + if self.__logger: + self.__logger.warning( + f"Invalid wildcard name '{escape_single_quotes(fullkey)}' in file '{escape_single_quotes(full_path)}'! (cannot start with underscore)" + ) else: self.wildcards[fullkey] = PPPWildcard(full_path, fullkey, choices) @@ -434,9 +457,10 @@ class PPPWildcards: with open(full_path, "r", encoding="utf-8") as file: content = yaml.safe_load(file) except: # pylint: disable=bare-except - self.__logger.warning( - f"Could not read file '{escape_single_quotes(full_path)}' with utf-8 encoding, trying windows-1252..." - ) + if self.__logger: + self.__logger.warning( + f"Could not read file '{escape_single_quotes(full_path)}' with utf-8 encoding, trying windows-1252..." + ) with open(full_path, "r", encoding="windows-1252") as file: content = yaml.safe_load(file) self.__add_wildcard(content, full_path, external_key_parts) @@ -455,9 +479,10 @@ class PPPWildcards: with open(full_path, "r", encoding="utf-8") as file: text_content = map(lambda x: x.strip("\n\r"), file.readlines()) except: # pylint: disable=bare-except - self.__logger.warning( - f"Could not read file '{escape_single_quotes(full_path)}' with utf-8 encoding, trying windows-1252..." - ) + if self.__logger: + self.__logger.warning( + f"Could not read file '{escape_single_quotes(full_path)}' with utf-8 encoding, trying windows-1252..." + ) with open(full_path, "r", encoding="windows-1252") as file: text_content = map(lambda x: x.strip("\n\r"), file.readlines()) text_content = list(filter(lambda x: x.strip() != "" and not x.strip().startswith("#"), text_content)) @@ -473,7 +498,8 @@ class PPPWildcards: directory (str): The path to the directory. """ if not os.path.exists(directory): - self.__logger.warning(f"Wildcard directory '{escape_single_quotes(directory)}' does not exist!") + if self.__logger: + self.__logger.warning(f"Wildcard directory '{escape_single_quotes(directory)}' does not exist!") return for filename in os.listdir(directory): full_path = os.path.abspath(os.path.join(directory, filename)) diff --git a/scripts/ppp_script.py b/scripts/ppp_script.py index 5868102..319617d 100644 --- a/scripts/ppp_script.py +++ b/scripts/ppp_script.py @@ -15,7 +15,7 @@ from modules.shared import opts # pylint: disable=import-error from modules.paths import models_path # pylint: disable=import-error import gradio as gr # pylint: disable=import-error from ppp import PromptPostProcessor -from ppp_classes import SUPPORTED_APPS, SUPPORTED_APPS_NAMES +from ppp_classes import IFWILDCARDS_CHOICES, ONWARNING_CHOICES, SUPPORTED_APPS, SUPPORTED_APPS_NAMES, PPPStateOptions from ppp_logging import DEBUG_LEVEL, PromptPostProcessorLogFactory from ppp_cache import PPPLRUCache from ppp_wildcards import PPPWildcards @@ -175,10 +175,54 @@ class PromptPostProcessorA1111Script(scripts.Script): ) ) ) + options = PPPStateOptions( + debug_level=DEBUG_LEVEL(getattr(opts, "ppp_gen_debug_level", PromptPostProcessor.DEFAULT_DEBUG_LEVEL)), + gen_onwarning=ONWARNING_CHOICES(getattr(opts, "ppp_gen_onwarning", PromptPostProcessor.DEFAULT_ONWARNING)), + wil_process_wildcards=getattr(opts, "ppp_wil_processwildcards", PromptPostProcessor.DEFAULT_WC_PROCESS), + wil_ifwildcards=IFWILDCARDS_CHOICES( + getattr(opts, "ppp_wil_ifwildcards", PromptPostProcessor.DEFAULT_IF_WILDCARDS) + ), + wil_choice_separator=getattr( + opts, "ppp_wil_choice_separator", PromptPostProcessor.DEFAULT_CHOICE_SEPARATOR + ), + wil_keep_choices_order=getattr( + opts, "ppp_wil_keep_choices_order", PromptPostProcessor.DEFAULT_KEEP_CHOICES_ORDER + ), + stn_separator=getattr(opts, "ppp_stn_separator", PromptPostProcessor.DEFAULT_STN_SEPARATOR), + stn_ignore_repeats=getattr(opts, "ppp_stn_ignorerepeats", PromptPostProcessor.DEFAULT_STN_IGNORE_REPEATS), + cup_do_cleanup=True, + cup_cleanup_variables=True, + cup_extraspaces=getattr(opts, "ppp_cup_extraspaces", PromptPostProcessor.DEFAULT_CUP_EXTRA_SPACES), + cup_emptyconstructs=getattr( + opts, "ppp_cup_emptyconstructs", PromptPostProcessor.DEFAULT_CUP_EMPTY_CONSTRUCTS + ), + cup_extraseparators=getattr( + opts, "ppp_cup_extraseparators", PromptPostProcessor.DEFAULT_CUP_EXTRA_SEPARATORS + ), + cup_extraseparators2=getattr( + opts, "ppp_cup_extraseparators2", PromptPostProcessor.DEFAULT_CUP_EXTRA_SEPARATORS2 + ), + cup_extraseparators_include_eol=getattr( + opts, + "ppp_cup_extraseparators_include_eol", + PromptPostProcessor.DEFAULT_CUP_EXTRA_SEPARATORS_INCLUDE_EOL, + ), + cup_breaks=getattr(opts, "ppp_cup_breaks", PromptPostProcessor.DEFAULT_CUP_BREAKS), + cup_breaks_eol=getattr(opts, "ppp_cup_breaks_eol", PromptPostProcessor.DEFAULT_CUP_BREAKS_EOL), + cup_ands=getattr(opts, "ppp_cup_ands", PromptPostProcessor.DEFAULT_CUP_ANDS), + cup_ands_eol=getattr(opts, "ppp_cup_ands_eol", PromptPostProcessor.DEFAULT_CUP_ANDS_EOL), + cup_extranetworktags=getattr( + opts, "ppp_cup_extranetworktags", PromptPostProcessor.DEFAULT_CUP_EXTRANETWORK_TAGS + ), + cup_mergeattention=getattr(opts, "ppp_cup_mergeattention", PromptPostProcessor.DEFAULT_CUP_MERGE_ATTENTION), + rem_removeextranetworktags=getattr( + opts, "ppp_rem_removeextranetworktags", PromptPostProcessor.DEFAULT_CUP_REMOVE_EXTRANETWORK_TAGS + ), + ) if self.ppp_logger is None: - lf = PromptPostProcessorLogFactory(app) + lf = PromptPostProcessorLogFactory() self.ppp_logger = lf.log - self.ppp_debug_level = DEBUG_LEVEL(getattr(opts, "ppp_gen_debug_level", DEBUG_LEVEL.none.value)) + self.ppp_debug_level = options.debug_level self.lru_cache = PPPLRUCache(1000, logger=self.ppp_logger, debug_level=self.ppp_debug_level) self.wildcards_obj = PPPWildcards(self.ppp_logger) self.extranetwork_mappings_obj = PPPExtraNetworkMappings(self.ppp_logger) @@ -190,7 +234,6 @@ class PromptPostProcessorA1111Script(scripts.Script): self.ppp_logger.warning("Compel parser is not supported!") init_images = getattr(p, "init_images", [None]) or [None] is_i2i = bool(init_images[0]) - self.ppp_debug_level = DEBUG_LEVEL(getattr(opts, "ppp_gen_debug_level", DEBUG_LEVEL.none.value)) do_i2i = getattr(opts, "ppp_gen_doi2i", False) add_prompts = getattr(opts, "ppp_gen_addpromptstometadata", True) if is_i2i and not do_i2i: @@ -237,60 +280,16 @@ class PromptPostProcessorA1111Script(scripts.Script): for f in en_mappings_folders.split(",") if f.strip() != "" ] - options = { - "debug_level": getattr(opts, "ppp_gen_debug_level", PromptPostProcessor.DEFAULT_DEBUG_LEVEL), - "on_warning": getattr(opts, "ppp_gen_onwarning", PromptPostProcessor.DEFAULT_ONWARNING), - "process_wildcards": getattr(opts, "ppp_wil_processwildcards", PromptPostProcessor.DEFAULT_WC_PROCESS), - "if_wildcards": getattr(opts, "ppp_wil_ifwildcards", PromptPostProcessor.DEFAULT_IF_WILDCARDS), - "choice_separator": getattr(opts, "ppp_wil_choice_separator", PromptPostProcessor.DEFAULT_CHOICE_SEPARATOR), - "keep_choices_order": getattr( - opts, "ppp_wil_keep_choices_order", PromptPostProcessor.DEFAULT_KEEP_CHOICES_ORDER - ), - "stn_separator": getattr(opts, "ppp_stn_separator", PromptPostProcessor.DEFAULT_STN_SEPARATOR), - "stn_ignore_repeats": getattr( - opts, "ppp_stn_ignorerepeats", PromptPostProcessor.DEFAULT_STN_IGNORE_REPEATS - ), - "do_cleanup": True, - "cleanup_variables": True, - "cleanup_extra_spaces": getattr(opts, "ppp_cup_extraspaces", PromptPostProcessor.DEFAULT_CUP_EXTRA_SPACES), - "cleanup_empty_constructs": getattr( - opts, "ppp_cup_emptyconstructs", PromptPostProcessor.DEFAULT_CUP_EMPTY_CONSTRUCTS - ), - "cleanup_extra_separators": getattr( - opts, "ppp_cup_extraseparators", PromptPostProcessor.DEFAULT_CUP_EXTRA_SEPARATORS - ), - "cleanup_extra_separators2": getattr( - opts, "ppp_cup_extraseparators2", PromptPostProcessor.DEFAULT_CUP_EXTRA_SEPARATORS2 - ), - "cleanup_extra_separators_include_eol": getattr( - opts, - "ppp_cup_extraseparators_include_eol", - PromptPostProcessor.DEFAULT_CUP_EXTRA_SEPARATORS_INCLUDE_EOL, - ), - "cleanup_breaks": getattr(opts, "ppp_cup_breaks", PromptPostProcessor.DEFAULT_CUP_BREAKS), - "cleanup_breaks_eol": getattr(opts, "ppp_cup_breaks_eol", PromptPostProcessor.DEFAULT_CUP_BREAKS_EOL), - "cleanup_ands": getattr(opts, "ppp_cup_ands", PromptPostProcessor.DEFAULT_CUP_ANDS), - "cleanup_ands_eol": getattr(opts, "ppp_cup_ands_eol", PromptPostProcessor.DEFAULT_CUP_ANDS_EOL), - "cleanup_extranetwork_tags": getattr( - opts, "ppp_cup_extranetworktags", PromptPostProcessor.DEFAULT_CUP_EXTRANETWORK_TAGS - ), - "cleanup_merge_attention": getattr( - opts, "ppp_cup_mergeattention", PromptPostProcessor.DEFAULT_CUP_MERGE_ATTENTION - ), - "remove_extranetwork_tags": getattr( - opts, "ppp_rem_removeextranetworktags", PromptPostProcessor.DEFAULT_CUP_REMOVE_EXTRANETWORK_TAGS - ), - } self.wildcards_obj.refresh_wildcards( - self.ppp_debug_level, wildcards_folders if options["process_wildcards"] else None + self.ppp_debug_level, wildcards_folders if options.wil_process_wildcards else None ) self.extranetwork_mappings_obj.refresh_extranetwork_mappings(self.ppp_debug_level, enmappings_folders) ppp = PromptPostProcessor( self.ppp_logger, - self.ppp_interrupt, env_info, options, self.grammar_content, + self.ppp_interrupt, self.wildcards_obj, self.extranetwork_mappings_obj, ) @@ -500,13 +499,13 @@ def on_ui_settings(): shared.opts.add_option( key="ppp_gen_onwarning", info=shared.OptionInfo( - default=PromptPostProcessor.ONWARNING_CHOICES.warn.value, + default=ONWARNING_CHOICES.warn.value, label="What to do on invalid content warnings?", component=gr.Radio, component_args={ "choices": ( - ("Show warning in console", PromptPostProcessor.ONWARNING_CHOICES.warn.value), - ("Stop the generation", PromptPostProcessor.ONWARNING_CHOICES.stop.value), + ("Show warning in console", ONWARNING_CHOICES.warn.value), + ("Stop the generation", ONWARNING_CHOICES.stop.value), ) }, section=section, @@ -567,16 +566,16 @@ def on_ui_settings(): info=shared.OptionInfo( default=import_old_settings( ["ppp_gen_ifwildcards", "ppp_ifwildcards"], - PromptPostProcessor.IFWILDCARDS_CHOICES.ignore.value, + IFWILDCARDS_CHOICES.ignore.value, ), label="What to do with remaining/invalid wildcards?", component=gr.Radio, component_args={ "choices": ( - ("Ignore", PromptPostProcessor.IFWILDCARDS_CHOICES.ignore.value), - ("Remove", PromptPostProcessor.IFWILDCARDS_CHOICES.remove.value), - ("Add visible warning", PromptPostProcessor.IFWILDCARDS_CHOICES.warn.value), - ("Stop the generation", PromptPostProcessor.IFWILDCARDS_CHOICES.stop.value), + ("Ignore", IFWILDCARDS_CHOICES.ignore.value), + ("Remove", IFWILDCARDS_CHOICES.remove.value), + ("Add visible warning", IFWILDCARDS_CHOICES.warn.value), + ("Stop the generation", IFWILDCARDS_CHOICES.stop.value), ) }, section=section, diff --git a/tests/base_tests.py b/tests/base_tests.py index b62296a..8ec08cd 100644 --- a/tests/base_tests.py +++ b/tests/base_tests.py @@ -1,9 +1,11 @@ +from dataclasses import replace import os import logging from typing import NamedTuple, Optional import unittest import datetime +from ppp_classes import IFWILDCARDS_CHOICES, ONWARNING_CHOICES, PPPStateOptions from ppp_enmappings import PPPExtraNetworkMappings # pylint: disable=import-error from ppp_wildcards import PPPWildcards # pylint: disable=import-error from ppp import PromptPostProcessor # pylint: disable=import-error @@ -36,35 +38,35 @@ class TestPromptPostProcessorBase(unittest.TestCase): else: log_filename = None # Disable file logging - self.lf = PromptPostProcessorLogFactory(None, log_filename) + self.lf = PromptPostProcessorLogFactory(log_filename) self.ppp_logger = self.lf.log self.ppp_logger.setLevel(logging.DEBUG) self.grammar_content = None self.interrupted = False - self.defopts = { - "debug_level": DEBUG_LEVEL.full.value, - "on_warning": PromptPostProcessor.ONWARNING_CHOICES.stop.value, - "process_wildcards": True, - "if_wildcards": PromptPostProcessor.IFWILDCARDS_CHOICES.ignore.value, - "choice_separator": ", ", - "keep_choices_order": False, - "stn_separator": ", ", - "stn_ignore_repeats": True, - "do_cleanup": True, - "cleanup_variables": True, - "cleanup_empty_constructs": True, - "cleanup_extra_separators": True, - "cleanup_extra_separators2": True, - "cleanup_extra_separators_include_eol": False, - "cleanup_extra_spaces": True, - "cleanup_breaks": True, - "cleanup_breaks_eol": False, - "cleanup_ands": True, - "cleanup_ands_eol": False, - "cleanup_extranetwork_tags": True, - "cleanup_merge_attention": True, - "remove_extranetwork_tags": False, - } + self.defopts = PPPStateOptions( + debug_level=DEBUG_LEVEL.full, + gen_onwarning=ONWARNING_CHOICES.stop, + wil_process_wildcards=True, + wil_ifwildcards=IFWILDCARDS_CHOICES.ignore, + wil_choice_separator=", ", + wil_keep_choices_order=False, + stn_separator=", ", + stn_ignore_repeats=True, + cup_do_cleanup=True, + cup_cleanup_variables=True, + cup_emptyconstructs=True, + cup_extraseparators=True, + cup_extraseparators2=True, + cup_extraseparators_include_eol=False, + cup_extraspaces=True, + cup_breaks=True, + cup_breaks_eol=False, + cup_ands=True, + cup_ands_eol=False, + cup_extranetworktags=True, + cup_mergeattention=True, + rem_removeextranetworktags=False, + ) self.def_env_info = { "app": "tests", "ppp_config": None, @@ -130,51 +132,37 @@ class TestPromptPostProcessorBase(unittest.TestCase): if ppp == "nocup": the_obj = PromptPostProcessor( self.ppp_logger, - self.interrupt, self.def_env_info, - { - **self.defopts, - "do_cleanup": False, - "cleanup_variables": False, - "cleanup_empty_constructs": False, - "cleanup_extra_separators": False, - "cleanup_extra_separators2": False, - "cleanup_extra_separators_include_eol": False, - "cleanup_extra_spaces": False, - "cleanup_breaks": False, - "cleanup_breaks_eol": False, - "cleanup_ands": False, - "cleanup_ands_eol": False, - "cleanup_extranetwork_tags": False, - "cleanup_merge_attention": False, - }, + replace( + self.defopts, + cup_do_cleanup=False, + cup_cleanup_variables=False, + cup_emptyconstructs=False, + cup_extraseparators=False, + cup_extraseparators2=False, + cup_extraseparators_include_eol=False, + cup_extraspaces=False, + cup_breaks=False, + cup_breaks_eol=False, + cup_ands=False, + cup_ands_eol=False, + cup_extranetworktags=False, + cup_mergeattention=False, + ), self.grammar_content, + self.interrupt, self.wildcards_obj, self.extranetwork_maps_obj, ) - # elif ppp == "comfyui": - # the_obj = PromptPostProcessor( - # self.ppp_logger, - # self.interrupt, - # { - # **self.def_env_info, - # "app": "comfyui", - # "model_class": "SDXL", - # }, - # self.defopts, - # self.grammar_content, - # self.wildcards_obj, - # self.extranetwork_maps_obj, - # ) else: the_obj = ppp if not the_obj: the_obj = PromptPostProcessor( self.ppp_logger, - self.interrupt, self.def_env_info, self.defopts, self.grammar_content, + self.interrupt, self.wildcards_obj, self.extranetwork_maps_obj, ) diff --git a/tests/tests_choices.py b/tests/tests_choices.py index fc4003a..4f553b6 100644 --- a/tests/tests_choices.py +++ b/tests/tests_choices.py @@ -1,3 +1,5 @@ +from dataclasses import replace + from ppp import PromptPostProcessor # pylint: disable=import-error from .base_tests import PromptPair, TestPromptPostProcessorBase @@ -81,10 +83,13 @@ class TestChoices(TestPromptPostProcessorBase): PromptPair("", ""), ppp=PromptPostProcessor( self.ppp_logger, - self.interrupt, self.def_env_info, - {**self.defopts, "remove_extranetwork_tags": True}, + replace( + self.defopts, + rem_removeextranetworktags=True, + ), self.grammar_content, + self.interrupt, self.wildcards_obj, self.extranetwork_maps_obj, ), diff --git a/tests/tests_cleanup.py b/tests/tests_cleanup.py index 4653435..abf0c6f 100644 --- a/tests/tests_cleanup.py +++ b/tests/tests_cleanup.py @@ -1,3 +1,5 @@ +from dataclasses import replace + from ppp import PromptPostProcessor # pylint: disable=import-error from .base_tests import PromptPair, TestPromptPostProcessorBase @@ -5,6 +7,7 @@ from .base_tests import PromptPair, TestPromptPostProcessorBase if __name__ == "__main__": raise SystemExit("This script must not be run directly") + class TestCleanup(TestPromptPostProcessorBase): def setUp(self): # pylint: disable=arguments-differ @@ -36,10 +39,13 @@ class TestCleanup(TestPromptPostProcessorBase): PromptPair("this is a test", ""), ppp=PromptPostProcessor( self.ppp_logger, - self.interrupt, self.def_env_info, - {**self.defopts, "remove_extranetwork_tags": True}, + replace( + self.defopts, + rem_removeextranetworktags=True, + ), self.grammar_content, + self.interrupt, self.wildcards_obj, self.extranetwork_maps_obj, ), @@ -51,14 +57,14 @@ class TestCleanup(TestPromptPostProcessorBase): PromptPair("this is a test,\nsecond line", ""), ppp=PromptPostProcessor( self.ppp_logger, - self.interrupt, self.def_env_info, - { - **self.defopts, - "cleanup_extra_separators2": False, - "cleanup_extra_separators_include_eol": False, - }, + replace( + self.defopts, + cup_extraseparators2=False, + cup_extraseparators_include_eol=False, + ), self.grammar_content, + self.interrupt, self.wildcards_obj, self.extranetwork_maps_obj, ), @@ -80,23 +86,23 @@ class TestCleanup(TestPromptPostProcessorBase): ), ppp=PromptPostProcessor( self.ppp_logger, - self.interrupt, self.def_env_info, - { - **self.defopts, - "cleanup_empty_constructs": False, - "cleanup_extra_separators": True, - "cleanup_extra_separators2": False, - "cleanup_extra_separators_include_eol": False, - "cleanup_extra_spaces": False, - "cleanup_breaks": False, - "cleanup_breaks_eol": False, - "cleanup_ands": False, - "cleanup_ands_eol": False, - "cleanup_extranetwork_tags": False, - "cleanup_merge_attention": False, - }, + replace( + self.defopts, + cup_emptyconstructs=False, + cup_extraseparators=True, + cup_extraseparators2=False, + cup_extraseparators_include_eol=False, + cup_extraspaces=False, + cup_breaks=False, + cup_breaks_eol=False, + cup_ands=False, + cup_ands_eol=False, + cup_extranetworktags=False, + cup_mergeattention=False, + ), self.grammar_content, + self.interrupt, self.wildcards_obj, self.extranetwork_maps_obj, ), diff --git a/tests/tests_commands.py b/tests/tests_commands.py index 110b7f1..b3d687a 100644 --- a/tests/tests_commands.py +++ b/tests/tests_commands.py @@ -45,13 +45,13 @@ class TestCommands(TestPromptPostProcessorBase): PromptPair("this is PONY", ""), ppp=PromptPostProcessor( self.ppp_logger, - self.interrupt, { **self.def_env_info, "model_filename": "./webui/models/Stable-diffusion/ponymodel.safetensors", }, self.defopts, self.grammar_content, + self.interrupt, self.wildcards_obj, self.extranetwork_maps_obj, ), @@ -269,13 +269,13 @@ class TestCommands(TestPromptPostProcessorBase): PromptPair("inlinetrigger, triggerpony1, triggerpony2", ""), ppp=PromptPostProcessor( self.ppp_logger, - self.interrupt, { **self.def_env_info, "model_filename": "./webui/models/Stable-diffusion/ponymodel.safetensors", }, self.defopts, self.grammar_content, + self.interrupt, self.wildcards_obj, self.extranetwork_maps_obj, ), @@ -290,13 +290,13 @@ class TestCommands(TestPromptPostProcessorBase): PromptPair("inlinetrigger, triggerpony1, triggerpony2", ""), ppp=PromptPostProcessor( self.ppp_logger, - self.interrupt, { **self.def_env_info, "model_filename": "./webui/models/Stable-diffusion/ponymodel.safetensors", }, self.defopts, self.grammar_content, + self.interrupt, self.wildcards_obj, self.extranetwork_maps_obj, ), @@ -311,13 +311,13 @@ class TestCommands(TestPromptPostProcessorBase): PromptPair("inlinetrigger, triggerpony1, triggerpony2", ""), ppp=PromptPostProcessor( self.ppp_logger, - self.interrupt, { **self.def_env_info, "model_filename": "./webui/models/Stable-diffusion/ponymodel.safetensors", }, self.defopts, self.grammar_content, + self.interrupt, self.wildcards_obj, self.extranetwork_maps_obj, ), @@ -332,13 +332,13 @@ class TestCommands(TestPromptPostProcessorBase): PromptPair("inlinetrigger, triggerillustrious1, triggerillustrious2", ""), ppp=PromptPostProcessor( self.ppp_logger, - self.interrupt, { **self.def_env_info, "model_filename": "./webui/models/Stable-diffusion/ilxlmodel.safetensors", }, self.defopts, self.grammar_content, + self.interrupt, self.wildcards_obj, self.extranetwork_maps_obj, ), diff --git a/tests/tests_host.py b/tests/tests_host.py index a04b801..f38ee98 100644 --- a/tests/tests_host.py +++ b/tests/tests_host.py @@ -21,13 +21,13 @@ class TestHosts(TestPromptPostProcessorBase): PromptPair("(test1:0.9) (test2) (test3:1.5) (test4:0.99)", ""), ppp=PromptPostProcessor( self.ppp_logger, - self.interrupt, { **self.def_env_info, "ppp_config": {"hosts": {"tests": {"attention": "parentheses"}}}, }, self.defopts, self.grammar_content, + self.interrupt, self.wildcards_obj, self.extranetwork_maps_obj, ), @@ -42,13 +42,13 @@ class TestHosts(TestPromptPostProcessorBase): PromptPair("test1 test2 test3", ""), ppp=PromptPostProcessor( self.ppp_logger, - self.interrupt, { **self.def_env_info, "ppp_config": {"hosts": {"tests": {"attention": "disable"}}}, }, self.defopts, self.grammar_content, + self.interrupt, self.wildcards_obj, self.extranetwork_maps_obj, ), @@ -63,13 +63,13 @@ class TestHosts(TestPromptPostProcessorBase): PromptPair("", ""), ppp=PromptPostProcessor( self.ppp_logger, - self.interrupt, { **self.def_env_info, "ppp_config": {"hosts": {"tests": {"attention": "remove"}}}, }, self.defopts, self.grammar_content, + self.interrupt, self.wildcards_obj, self.extranetwork_maps_obj, ), @@ -84,13 +84,13 @@ class TestHosts(TestPromptPostProcessorBase): PromptPair("", ""), ppp=PromptPostProcessor( self.ppp_logger, - self.interrupt, { **self.def_env_info, "ppp_config": {"hosts": {"tests": {"attention": "error"}}}, }, self.defopts, self.grammar_content, + self.interrupt, self.wildcards_obj, self.extranetwork_maps_obj, ), @@ -106,13 +106,13 @@ class TestHosts(TestPromptPostProcessorBase): PromptPair("test1", ""), ppp=PromptPostProcessor( self.ppp_logger, - self.interrupt, { **self.def_env_info, "ppp_config": {"hosts": {"tests": {"scheduling": "before"}}}, }, self.defopts, self.grammar_content, + self.interrupt, self.wildcards_obj, self.extranetwork_maps_obj, ), @@ -127,13 +127,13 @@ class TestHosts(TestPromptPostProcessorBase): PromptPair("test2", ""), ppp=PromptPostProcessor( self.ppp_logger, - self.interrupt, { **self.def_env_info, "ppp_config": {"hosts": {"tests": {"scheduling": "after"}}}, }, self.defopts, self.grammar_content, + self.interrupt, self.wildcards_obj, self.extranetwork_maps_obj, ), @@ -148,13 +148,13 @@ class TestHosts(TestPromptPostProcessorBase): PromptPair("test1 test3", ""), ppp=PromptPostProcessor( self.ppp_logger, - self.interrupt, { **self.def_env_info, "ppp_config": {"hosts": {"tests": {"scheduling": "first"}}}, }, self.defopts, self.grammar_content, + self.interrupt, self.wildcards_obj, self.extranetwork_maps_obj, ), @@ -169,13 +169,13 @@ class TestHosts(TestPromptPostProcessorBase): PromptPair("", ""), ppp=PromptPostProcessor( self.ppp_logger, - self.interrupt, { **self.def_env_info, "ppp_config": {"hosts": {"tests": {"scheduling": "remove"}}}, }, self.defopts, self.grammar_content, + self.interrupt, self.wildcards_obj, self.extranetwork_maps_obj, ), @@ -190,13 +190,13 @@ class TestHosts(TestPromptPostProcessorBase): PromptPair("", ""), ppp=PromptPostProcessor( self.ppp_logger, - self.interrupt, { **self.def_env_info, "ppp_config": {"hosts": {"tests": {"scheduling": "error"}}}, }, self.defopts, self.grammar_content, + self.interrupt, self.wildcards_obj, self.extranetwork_maps_obj, ), @@ -212,13 +212,13 @@ class TestHosts(TestPromptPostProcessorBase): PromptPair("test1", ""), ppp=PromptPostProcessor( self.ppp_logger, - self.interrupt, { **self.def_env_info, "ppp_config": {"hosts": {"tests": {"alternation": "first"}}}, }, self.defopts, self.grammar_content, + self.interrupt, self.wildcards_obj, self.extranetwork_maps_obj, ), @@ -233,13 +233,13 @@ class TestHosts(TestPromptPostProcessorBase): PromptPair("", ""), ppp=PromptPostProcessor( self.ppp_logger, - self.interrupt, { **self.def_env_info, "ppp_config": {"hosts": {"tests": {"alternation": "remove"}}}, }, self.defopts, self.grammar_content, + self.interrupt, self.wildcards_obj, self.extranetwork_maps_obj, ), @@ -254,13 +254,13 @@ class TestHosts(TestPromptPostProcessorBase): PromptPair("", ""), ppp=PromptPostProcessor( self.ppp_logger, - self.interrupt, { **self.def_env_info, "ppp_config": {"hosts": {"tests": {"alternation": "error"}}}, }, self.defopts, self.grammar_content, + self.interrupt, self.wildcards_obj, self.extranetwork_maps_obj, ), @@ -276,13 +276,13 @@ class TestHosts(TestPromptPostProcessorBase): PromptPair("test1\ntest2", ""), ppp=PromptPostProcessor( self.ppp_logger, - self.interrupt, { **self.def_env_info, "ppp_config": {"hosts": {"tests": {"and": "eol"}}}, }, self.defopts, self.grammar_content, + self.interrupt, self.wildcards_obj, self.extranetwork_maps_obj, ), @@ -297,13 +297,13 @@ class TestHosts(TestPromptPostProcessorBase): PromptPair("test1, test2", ""), ppp=PromptPostProcessor( self.ppp_logger, - self.interrupt, { **self.def_env_info, "ppp_config": {"hosts": {"tests": {"and": "comma"}}}, }, self.defopts, self.grammar_content, + self.interrupt, self.wildcards_obj, self.extranetwork_maps_obj, ), @@ -318,13 +318,13 @@ class TestHosts(TestPromptPostProcessorBase): PromptPair("test1 test2", ""), ppp=PromptPostProcessor( self.ppp_logger, - self.interrupt, { **self.def_env_info, "ppp_config": {"hosts": {"tests": {"and": "remove"}}}, }, self.defopts, self.grammar_content, + self.interrupt, self.wildcards_obj, self.extranetwork_maps_obj, ), @@ -339,13 +339,13 @@ class TestHosts(TestPromptPostProcessorBase): PromptPair("", ""), ppp=PromptPostProcessor( self.ppp_logger, - self.interrupt, { **self.def_env_info, "ppp_config": {"hosts": {"tests": {"and": "error"}}}, }, self.defopts, self.grammar_content, + self.interrupt, self.wildcards_obj, self.extranetwork_maps_obj, ), @@ -361,13 +361,13 @@ class TestHosts(TestPromptPostProcessorBase): PromptPair("test1\ntest2", ""), ppp=PromptPostProcessor( self.ppp_logger, - self.interrupt, { **self.def_env_info, "ppp_config": {"hosts": {"tests": {"break": "eol"}}}, }, self.defopts, self.grammar_content, + self.interrupt, self.wildcards_obj, self.extranetwork_maps_obj, ), @@ -382,13 +382,13 @@ class TestHosts(TestPromptPostProcessorBase): PromptPair("test1, test2", ""), ppp=PromptPostProcessor( self.ppp_logger, - self.interrupt, { **self.def_env_info, "ppp_config": {"hosts": {"tests": {"break": "comma"}}}, }, self.defopts, self.grammar_content, + self.interrupt, self.wildcards_obj, self.extranetwork_maps_obj, ), @@ -403,13 +403,13 @@ class TestHosts(TestPromptPostProcessorBase): PromptPair("test1 test2", ""), ppp=PromptPostProcessor( self.ppp_logger, - self.interrupt, { **self.def_env_info, "ppp_config": {"hosts": {"tests": {"break": "remove"}}}, }, self.defopts, self.grammar_content, + self.interrupt, self.wildcards_obj, self.extranetwork_maps_obj, ), @@ -424,13 +424,13 @@ class TestHosts(TestPromptPostProcessorBase): PromptPair("", ""), ppp=PromptPostProcessor( self.ppp_logger, - self.interrupt, { **self.def_env_info, "ppp_config": {"hosts": {"tests": {"break": "error"}}}, }, self.defopts, self.grammar_content, + self.interrupt, self.wildcards_obj, self.extranetwork_maps_obj, ), diff --git a/tests/tests_variables.py b/tests/tests_variables.py index aafce6e..37e2ab9 100644 --- a/tests/tests_variables.py +++ b/tests/tests_variables.py @@ -1,4 +1,7 @@ -from ppp import PromptPostProcessor # pylint: disable=import-error +from dataclasses import replace + +from ppp import PromptPostProcessor +from ppp_classes import ONWARNING_CHOICES # pylint: disable=import-error from .base_tests import PromptPair, TestPromptPostProcessorBase if __name__ == "__main__": @@ -46,10 +49,13 @@ class TestVariables(TestPromptPostProcessorBase): variables={"v1": ""}, ppp=PromptPostProcessor( self.ppp_logger, - self.interrupt, self.def_env_info, - {**self.defopts, "on_warning": PromptPostProcessor.ONWARNING_CHOICES.warn.value}, + replace( + self.defopts, + gen_onwarning=ONWARNING_CHOICES.warn, + ), self.grammar_content, + self.interrupt, self.wildcards_obj, self.extranetwork_maps_obj, ), @@ -147,10 +153,13 @@ class TestVariables(TestPromptPostProcessorBase): PromptPair("NO", ""), ppp=PromptPostProcessor( self.ppp_logger, - self.interrupt, self.def_env_info, - {**self.defopts, "on_warning": PromptPostProcessor.ONWARNING_CHOICES.warn.value}, + replace( + self.defopts, + gen_onwarning=ONWARNING_CHOICES.warn, + ), self.grammar_content, + self.interrupt, self.wildcards_obj, self.extranetwork_maps_obj, ), @@ -175,10 +184,13 @@ class TestVariables(TestPromptPostProcessorBase): PromptPair("NO", ""), ppp=PromptPostProcessor( self.ppp_logger, - self.interrupt, self.def_env_info, - {**self.defopts, "on_warning": PromptPostProcessor.ONWARNING_CHOICES.warn.value}, + replace( + self.defopts, + gen_onwarning=ONWARNING_CHOICES.warn, + ), self.grammar_content, + self.interrupt, self.wildcards_obj, self.extranetwork_maps_obj, ), @@ -203,10 +215,13 @@ class TestVariables(TestPromptPostProcessorBase): PromptPair("NO", ""), ppp=PromptPostProcessor( self.ppp_logger, - self.interrupt, self.def_env_info, - {**self.defopts, "on_warning": PromptPostProcessor.ONWARNING_CHOICES.warn.value}, + replace( + self.defopts, + gen_onwarning=ONWARNING_CHOICES.warn, + ), self.grammar_content, + self.interrupt, self.wildcards_obj, self.extranetwork_maps_obj, ), diff --git a/tests/tests_variants.py b/tests/tests_variants.py index 2e7ef26..80115b0 100644 --- a/tests/tests_variants.py +++ b/tests/tests_variants.py @@ -1,4 +1,7 @@ -from ppp import PromptPostProcessor # pylint: disable=import-error +from dataclasses import replace + +from ppp import PromptPostProcessor +from ppp_classes import ONWARNING_CHOICES # pylint: disable=import-error from .base_tests import PromptPair, TestPromptPostProcessorBase if __name__ == "__main__": @@ -21,7 +24,6 @@ class TestModelVariants(TestPromptPostProcessorBase): PromptPair("test1test2", ""), ppp=PromptPostProcessor( self.ppp_logger, - self.interrupt, { **self.def_env_info, "model_filename": "./webui/models/Stable-diffusion/testmodel.safetensors", @@ -61,11 +63,12 @@ class TestModelVariants(TestPromptPostProcessorBase): } }, }, - { - **self.defopts, - "on_warning": PromptPostProcessor.ONWARNING_CHOICES.warn.value, - }, + replace( + self.defopts, + gen_onwarning=ONWARNING_CHOICES.warn, + ), self.grammar_content, + self.interrupt, self.wildcards_obj, self.extranetwork_maps_obj, ), diff --git a/tests/tests_wildcards.py b/tests/tests_wildcards.py index 1cf041e..3af4362 100644 --- a/tests/tests_wildcards.py +++ b/tests/tests_wildcards.py @@ -1,4 +1,7 @@ +from dataclasses import replace + from ppp import PromptPostProcessor +from ppp_classes import IFWILDCARDS_CHOICES from .base_tests import PromptPair, TestPromptPostProcessorBase if __name__ == "__main__": @@ -18,14 +21,14 @@ class TestWildcards(TestPromptPostProcessorBase): PromptPair("__bad_wildcard__", "{option1|option2}"), ppp=PromptPostProcessor( self.ppp_logger, - self.interrupt, self.def_env_info, - { - **self.defopts, - "process_wildcards": False, - "if_wildcards": PromptPostProcessor.IFWILDCARDS_CHOICES.ignore.value, - }, + replace( + self.defopts, + wil_process_wildcards=False, + wil_ifwildcards=IFWILDCARDS_CHOICES.ignore, + ), self.grammar_content, + self.interrupt, self.wildcards_obj, self.extranetwork_maps_obj, ), @@ -43,14 +46,14 @@ class TestWildcards(TestPromptPostProcessorBase): ), ppp=PromptPostProcessor( self.ppp_logger, - self.interrupt, self.def_env_info, - { - **self.defopts, - "process_wildcards": False, - "if_wildcards": PromptPostProcessor.IFWILDCARDS_CHOICES.remove.value, - }, + replace( + self.defopts, + wil_process_wildcards=False, + wil_ifwildcards=IFWILDCARDS_CHOICES.remove, + ), self.grammar_content, + self.interrupt, self.wildcards_obj, self.extranetwork_maps_obj, ), @@ -62,14 +65,14 @@ class TestWildcards(TestPromptPostProcessorBase): PromptPair(PromptPostProcessor.WILDCARD_WARNING + "__bad_wildcard__", "{option1|option2}"), ppp=PromptPostProcessor( self.ppp_logger, - self.interrupt, self.def_env_info, - { - **self.defopts, - "process_wildcards": False, - "if_wildcards": PromptPostProcessor.IFWILDCARDS_CHOICES.warn.value, - }, + replace( + self.defopts, + wil_process_wildcards=False, + wil_ifwildcards=IFWILDCARDS_CHOICES.warn, + ), self.grammar_content, + self.interrupt, self.wildcards_obj, self.extranetwork_maps_obj, ), @@ -84,14 +87,14 @@ class TestWildcards(TestPromptPostProcessorBase): ), ppp=PromptPostProcessor( self.ppp_logger, - self.interrupt, self.def_env_info, - { - **self.defopts, - "process_wildcards": False, - "if_wildcards": PromptPostProcessor.IFWILDCARDS_CHOICES.stop.value, - }, + replace( + self.defopts, + wil_process_wildcards=False, + wil_ifwildcards=IFWILDCARDS_CHOICES.stop, + ), self.grammar_content, + self.interrupt, self.wildcards_obj, self.extranetwork_maps_obj, ), @@ -104,14 +107,14 @@ class TestWildcards(TestPromptPostProcessorBase): PromptPair(PromptPostProcessor.WILDCARD_WARNING + "__bad_wildcard__", ""), ppp=PromptPostProcessor( self.ppp_logger, - self.interrupt, self.def_env_info, - { - **self.defopts, - "process_wildcards": False, - "if_wildcards": PromptPostProcessor.IFWILDCARDS_CHOICES.warn.value, - }, + replace( + self.defopts, + wil_process_wildcards=False, + wil_ifwildcards=IFWILDCARDS_CHOICES.warn, + ), self.grammar_content, + self.interrupt, self.wildcards_obj, self.extranetwork_maps_obj, ),