From 47375db1ba76d0b2a2644c67a90d99957bc20168 Mon Sep 17 00:00:00 2001 From: Antonio Cordero Balcazar Date: Sat, 4 Apr 2026 19:04:46 +0200 Subject: [PATCH] * Escape quotes on log messages * Separated tests in multiple files. * Added pydantic validation for configuration. * Improved backup/restore on visit of nodes. * Echoing of variables with default no longer visits the default when not used (and no longer creates variables that were not actually used). --- .vscode/launch.json | 8 - docs/SYNTAX.md | 18 +- ppp.py | 198 ++-- ppp_cache.py | 2 +- ppp_classes.py | 151 +++ ppp_comfyui.py | 18 +- ppp_enmappings.py | 28 +- ppp_hosts.py | 19 - ppp_logging.py | 2 +- ppp_utils.py | 24 + ppp_wildcards.py | 44 +- scripts/ppp_script.py | 14 +- tests/base_tests.py | 207 ++++ tests/tests.py | 1961 ------------------------------------ tests/tests_choices.py | 101 ++ tests/tests_cleanup.py | 132 +++ tests/tests_commands.py | 348 +++++++ tests/tests_host.py | 442 ++++++++ tests/tests_performance.py | 78 ++ tests/tests_stn.py | 127 +++ tests/tests_variables.py | 160 +++ tests/tests_variants.py | 75 ++ tests/tests_wildcards.py | 400 ++++++++ 23 files changed, 2440 insertions(+), 2117 deletions(-) create mode 100644 ppp_classes.py delete mode 100644 ppp_hosts.py create mode 100644 tests/base_tests.py delete mode 100644 tests/tests.py create mode 100644 tests/tests_choices.py create mode 100644 tests/tests_cleanup.py create mode 100644 tests/tests_commands.py create mode 100644 tests/tests_host.py create mode 100644 tests/tests_performance.py create mode 100644 tests/tests_stn.py create mode 100644 tests/tests_variables.py create mode 100644 tests/tests_variants.py create mode 100644 tests/tests_wildcards.py diff --git a/.vscode/launch.json b/.vscode/launch.json index 264f4c8..4cb3ae9 100644 --- a/.vscode/launch.json +++ b/.vscode/launch.json @@ -9,14 +9,6 @@ "type": "debugpy", "request": "attach", "processId": "${command:pickProcess}" - }, - { - "name": "Tests", - "type": "debugpy", - "request": "launch", - "program": "tests/tests.py", - "console": "integratedTerminal", - "justMyCode": true } ] } \ No newline at end of file diff --git a/docs/SYNTAX.md b/docs/SYNTAX.md index d986f85..7c9dc71 100644 --- a/docs/SYNTAX.md +++ b/docs/SYNTAX.md @@ -60,15 +60,15 @@ The only command available is `include wildcard`, which will include the choices These are examples of formats you can use to insert a choice construct: -| Construct | Result | -| --------- | ------ | -| `{choice1\|5::choice2\|3::choice3}` | select 1 choice, two of them have weights | -| `{3$$choice1\|5 if _is_sd1::choice2\|choice3}` | select 3 choices, one has a weight and a condition | -| `{2-3$$2::choice1\|choice2\|choice3}` | select 2 to 3 choices, one of them has a weight | -| `{r2-3$$choice1\|choice2\|choice3}` | select 2 to 3 choices allowing repetition | -| `{2-3$$ / $$choice1\|choice2\|choice3}` | select 2 to 3 choices with separator " / " | -| `{o$$if _is_sd1::choice1\|if _is_sd2::choice2}`| select 1 choice, both have conditions, if none matches it is allowed because we indicate that it is optional | -| `{choice1\|choice2\|%0.5::path/wildcard}` | select 1 choice from the two specified and the ones inside the path/wildcard wildcard, which will be weighted with half their weights | +| Construct | Result | +| --------- | ------ | +| `{choice1\|5::choice2\|3::choice3}` | select 1 choice, two of them have weights | +| `{3$$choice1\|5 if _is_sd1::choice2\|choice3}` | select 3 choices, one has a weight and a condition | +| `{2-3$$2::choice1\|choice2\|choice3}` | select 2 to 3 choices, one of them has a weight | +| `{r2-3$$choice1\|choice2\|choice3}` | select 2 to 3 choices allowing repetition | +| `{2-3$$ / $$choice1\|choice2\|choice3}` | select 2 to 3 choices with separator " / " | +| `{o$$if _is_sd1::choice1\|if _is_sd2::choice2}` | select 1 choice, both have conditions, if none matches it is allowed because we indicate that it is optional | +| `{choice1\|choice2\|%0.5::include path/wildcard}` | select 1 choice from the two specified and the ones inside the path/wildcard wildcard, which will be weighted with half their weights | Notes: diff --git a/ppp.py b/ppp.py index 63109f1..ea3e59c 100644 --- a/ppp.py +++ b/ppp.py @@ -12,10 +12,11 @@ import lark import numpy as np import yaml -from ppp_hosts import SUPPORTED_APPS # pylint: disable=import-error -from ppp_logging import DEBUG_LEVEL # pylint: disable=import-error -from ppp_wildcards import PPPWildcard, PPPWildcards # pylint: disable=import-error -from ppp_enmappings import PPPENMappingVariant, PPPExtraNetworkMappings # pylint: disable=import-error +from ppp_classes import SUPPORTED_APPS +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): @@ -132,14 +133,15 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in self.config: dict[str, Any] = yaml.safe_load(f) except Exception as exc: # pylint: disable=broad-exception-caught self.config = {} - raise PPPInterrupt(f"Failed to load default configuration from '{default_config_file}'.") from exc + raise PPPInterrupt( + f"Failed to load default configuration from '{escape_single_quotes(default_config_file)}'." + ) from exc validate_def_cfg = self.__validate_normalize_configuration(self.config, "default configuration file") if validate_def_cfg != 0: errmsg = "Default configuration file has errors. Please restore the default configuration file and, per instructions, use a copy to adapt it." if validate_def_cfg == 2: raise PPPInterrupt(errmsg) - else: - self.logger.warning(errmsg) + self.logger.warning(errmsg) user_config_file = self.env_info.get("ppp_config", "") user_config: dict[str, Any] = {} @@ -162,8 +164,9 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in if user_config_file and os.path.exists(user_config_file): with open(user_config_file, "r", encoding="utf-8") as f: user_config = yaml.safe_load(f) - self.__validate_normalize_configuration(user_config, "user configuration") - self.__merge_configuration(user_config) + self.__validate_normalize_configuration(user_config, "user configuration") + if user_config: + self.__merge_configuration(user_config) self.models_config: dict[str, dict[str, Any] | None] = self.config.get("models") or {} self.known_models: list[str] = list(self.models_config.keys()) @@ -179,7 +182,7 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in self.host_config: dict[str, Any] = (self.config.get("hosts") or {}).get(self.env_info.get("app", "")) if self.host_config is None: raise PPPInterrupt( - f"No host configuration found for app '{self.env_info.get('app', '')}'. Please check your configuration." + f"No host configuration found for app '{escape_single_quotes(self.env_info.get('app', ''))}'. Please check your configuration." ) # Update env_info with model detection @@ -208,7 +211,7 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in self.variants_definitions[v] = (m, vo["find_in_filename"]) else: self.logger.warning( - f"Variant name '{v}' in model '{m}' conflicts with a known model name. Discarding variant." + f"Variant name '{escape_single_quotes(v)}' in model '{escape_single_quotes(m)}' conflicts with a known model name. Discarding variant." ) if self.debug_level != DEBUG_LEVEL.none: @@ -442,27 +445,27 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in return {"regex": find_in_filename, "flags": re.IGNORECASE} except re.error: self.logger.warning( - f"{where.title()}: Invalid regex pattern for variant '{variant_key}' in model '{model_key}'. Discarding variant." + f"{where.title()}: Invalid regex pattern for variant '{escape_single_quotes(variant_key)}' in model '{escape_single_quotes(model_key)}'. Discarding variant." ) elif isinstance(find_in_filename, dict): regex = find_in_filename.get("regex", "") flags = find_in_filename.get("flags", []) if not isinstance(regex, str) or not isinstance(flags, list) or not all(isinstance(f, str) for f in flags): self.logger.warning( - f"{where.title()}: Invalid format for 'find_in_filename' for variant '{variant_key}' in model '{model_key}'. Discarding variant." + f"{where.title()}: Invalid format for 'find_in_filename' for variant '{escape_single_quotes(variant_key)}' in model '{escape_single_quotes(model_key)}'. Discarding variant." ) else: fl = self.__re_flags_from_list(flags) if fl == 0 and len(flags): self.logger.warning( - f"{where.title()}: Invalid regex flags for variant '{variant_key}' in model '{model_key}'. Discarding variant." + f"{where.title()}: Invalid regex flags for variant '{escape_single_quotes(variant_key)}' in model '{escape_single_quotes(model_key)}'. Discarding variant." ) try: re.compile(regex, fl) return {"regex": regex, "flags": fl} except re.error: self.logger.warning( - f"{where.title()}: Invalid regex pattern for variant '{variant_key}' in model '{model_key}'. Discarding variant." + f"{where.title()}: Invalid regex pattern for variant '{escape_single_quotes(variant_key)}' in model '{escape_single_quotes(model_key)}'. Discarding variant." ) return None @@ -493,14 +496,18 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in defcfg_hosts: dict[str, Any] = cfg.get("hosts", {}) for host_key, host_value in dict(defcfg_hosts).items(): if host_key not in SUPPORTED_APPS._value2member_map_: # pylint: disable=protected-access - self.logger.warning(f"{where.capitalize()}: Unsupported host '{host_key}'. Discarding host.") + self.logger.warning( + f"{where.capitalize()}: Unsupported host '{escape_single_quotes(host_key)}'. Discarding host." + ) defcfg_hosts.pop(host_key, None) result = 1 elif host_value is not None and ( not isinstance(host_value, dict) or not all(k in ["attention", "scheduling", "alternation", "and", "break"] for k in host_value) ): - self.logger.warning(f"{where.capitalize()}: Invalid format for host '{host_key}'. Discarding host.") + self.logger.warning( + f"{where.capitalize()}: Invalid format for host '{escape_single_quotes(host_key)}'. Discarding host." + ) defcfg_hosts.pop(host_key, None) result = 1 defcfg_models: dict[str, Any] = cfg.get("models", {}) @@ -510,7 +517,9 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in or model_value.get("detect") is None or not isinstance(model_value["detect"], dict) ): - self.logger.warning(f"{where.capitalize()}: Invalid format for model '{model_key}'. Discarding model.") + self.logger.warning( + f"{where.capitalize()}: Invalid format for model '{escape_single_quotes(model_key)}'. Discarding model." + ) defcfg_models.pop(model_key, None) result = 1 else: @@ -518,14 +527,14 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in for host_key, host_value in dict(defcfg_m_detect).items(): if host_key not in SUPPORTED_APPS._value2member_map_: # pylint: disable=protected-access self.logger.warning( - f"{where.capitalize()}: Unsupported host '{host_key}' in 'detect' for model '{model_key}'. Discarding host." + f"{where.capitalize()}: Unsupported host '{escape_single_quotes(host_key)}' in 'detect' for model '{escape_single_quotes(model_key)}'. Discarding host." ) defcfg_m_detect.pop(host_key, None) result = 1 elif host_value is not None: if not isinstance(host_value, dict): self.logger.warning( - f"{where.capitalize()}: Invalid format for host '{host_key}' in 'detect' for model '{model_key}'. Discarding host." + f"{where.capitalize()}: Invalid format for host '{escape_single_quotes(host_key)}' in 'detect' for model '{escape_single_quotes(model_key)}'. Discarding host." ) defcfg_m_detect.pop(host_key, None) result = 1 @@ -534,27 +543,27 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in isinstance(c, str) for c in host_value["class"] ): self.logger.warning( - f"{where.capitalize()}: Invalid format for 'class' in host '{host_key}' in 'detect' for model '{model_key}'. Discarding host." + f"{where.capitalize()}: Invalid format for 'class' in host '{escape_single_quotes(host_key)}' in 'detect' for model '{escape_single_quotes(model_key)}'. Discarding host." ) defcfg_m_detect.pop(host_key, None) result = 1 elif "property" in host_value: if not isinstance(host_value["property"], str): self.logger.warning( - f"{where.capitalize()}: Invalid format for 'property' in host '{host_key}' in 'detect' for model '{model_key}'. Discarding host." + f"{where.capitalize()}: Invalid format for 'property' in host '{escape_single_quotes(host_key)}' in 'detect' for model '{escape_single_quotes(model_key)}'. Discarding host." ) defcfg_m_detect.pop(host_key, None) result = 1 else: self.logger.warning( - f"{where.capitalize()}: Neither 'class' nor 'property' specified for host '{host_key}' in 'detect' for model '{model_key}'. Discarding host." + f"{where.capitalize()}: Neither 'class' nor 'property' specified for host '{escape_single_quotes(host_key)}' in 'detect' for model '{escape_single_quotes(model_key)}'. Discarding host." ) defcfg_m_detect.pop(host_key, None) result = 1 if "variants" in model_value: if not isinstance(model_value["variants"], dict): self.logger.warning( - f"{where.capitalize()}: Invalid format for 'variants' in model '{model_key}'. Discarding model." + f"{where.capitalize()}: Invalid format for 'variants' in model '{escape_single_quotes(model_key)}'. Discarding model." ) defcfg_models.pop(model_key, None) result = 1 @@ -563,7 +572,7 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in for variant_key, variant_value in dict(defcfg_m_variants).items(): if not isinstance(variant_key, str) or not variant_key.isidentifier(): self.logger.warning( - f"{where.capitalize()}: Invalid variant name '{variant_key}' in model '{model_key}'. Discarding variant." + f"{where.capitalize()}: Invalid variant name '{escape_single_quotes(variant_key)}' in model '{escape_single_quotes(model_key)}'. Discarding variant." ) defcfg_m_variants.pop(variant_key, None) result = 1 @@ -571,7 +580,7 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in variant_value.get("find_in_filename"), (str, dict, list) ): self.logger.warning( - f"{where.capitalize()}: Invalid format for variant '{variant_key}' in model '{model_key}'. Discarding variant." + f"{where.capitalize()}: Invalid format for variant '{escape_single_quotes(variant_key)}' in model '{escape_single_quotes(model_key)}'. Discarding variant." ) defcfg_m_variants.pop(variant_key, None) result = 1 @@ -1193,7 +1202,7 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in parsed_prompt = None try: if self.debug_level == DEBUG_LEVEL.full: - self.logger.debug(self.format_output(f"Parsing {prompt_description}: '{prompt}'")) + 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): @@ -1206,7 +1215,7 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in except lark.exceptions.UnexpectedInput: if raise_parsing_error: raise - self.logger.exception(self.format_output(f"Parsing failed on prompt!: {prompt}")) + 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") @@ -1306,13 +1315,19 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in 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 = {k: v.copy() for k, v in self.__ppp.user_variables.items()} + backup_echoed_variables = self.__ppp.echoed_variables.copy() if node is not None: if isinstance(node, list): for child in node: @@ -1329,12 +1344,16 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in 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: @@ -1416,7 +1435,7 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in info = f"({info}) " if info is not None and info != "" else "" output = self.result[len(start_result) :] if output != "": - output = f" >> '{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}" @@ -1512,7 +1531,9 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in try: var_value_adjusted = int(var_value) except (ValueError, TypeError): - self.warn_or_stop(f"Cannot convert variable value '{var_value}' to integer for comparison") + 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: @@ -1820,7 +1841,7 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in self.__negtags.append( self.NegTag(len(self.result), len(self.result), content, parameters, self.__shell.copy()) ) - info = f"with {parameters or 'no parameters'} : {content}" + 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::]) @@ -1862,18 +1883,20 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in t1 = time.monotonic_ns() start_result = self.result if variable.startswith("_"): - self.warn_or_stop(f"Invalid variable name '{variable}' detected! System variables cannot be set.") + 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] = [m.value for m in modifiers.children] if modifiers is not None else [] if any(item in modifiers_str for item in ["+", "add"]): - info += f" += '{value_description}'" + 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 {variable}") + self.warn_or_stop(f"Unknown variable {escape_single_quotes(variable)}") elif isinstance(raw_oldvalue, str): newvalue = lark.Tree( lark.Token("RULE", "varvalue"), @@ -1887,7 +1910,7 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in # Meta should be {"content": raw_oldvalue.meta.content + value.meta.content}, ) elif any(item in modifiers_str for item in ["?", "ifundefined"]): - info += f" ?= '{value_description}'" + info += f" ?= '{escape_single_quotes(value_description or '')}'" raw_oldvalue = self.__ppp.user_variables.get(variable, None) if raw_oldvalue is None: newvalue = value @@ -1907,7 +1930,7 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in if currentvalue is None: info += "not evaluated yet" else: - info += f"'{currentvalue}'" + info += f"'{escape_single_quotes(currentvalue)}'" t2 = time.monotonic_ns() self.__debug_end(command, start_result, t2 - t1, info) @@ -1934,22 +1957,28 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in """ t1 = time.monotonic_ns() start_result = self.result - if default is not None: - default_value = self.__visit(default, True) # for log + 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) + default_value = v self.__ppp.echoed_variables[variable] = v self.result += v else: - self.warn_or_stop(f"Unknown variable {variable}") + self.warn_or_stop(f"Unknown variable {escape_single_quotes(variable)}") else: self.__ppp.echoed_variables[variable] = value t2 = time.monotonic_ns() info = variable - if default is not None: - info += f" with default '{default_value}'" + 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): @@ -2039,7 +2068,7 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in ) except lark.exceptions.UnexpectedInput as e: self.warn_or_stop( - f"Error parsing condition '{v.condition}' in extranetwork mapping '{extnet_id}'! : {e.__class__.__name__}", + f"Error parsing condition '{escape_single_quotes(v.condition)}' in extranetwork mapping '{escape_single_quotes(extnet_id)}'! : {e.__class__.__name__}", e, ) cnd = None @@ -2064,7 +2093,7 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in if found.name: if not found_in_cache and self.__ppp.debug_level != DEBUG_LEVEL.none: self.__ppp.logger.info( - f"Mapping extranetwork '{extnet_id}' to '{extnet_type}:{found.name}'" + 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 @@ -2083,11 +2112,15 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in 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 '{extnet_id}' to just triggers") + 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 '{extnet_id}' to nothing") + 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) @@ -2097,12 +2130,12 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in ) except lark.exceptions.UnexpectedInput as e: self.warn_or_stop( - f"Error parsing triggers '{extra_triggers}' in extranetwork mapping '{extnet_id}'! : {e.__class__.__name__}", + 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 '{extnet_id}' not found!") + 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 @@ -2134,13 +2167,15 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in 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 '{wildcard_key}' not found for default filter setting!") + 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 '{wc}'") + 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: list[list[str]] = [ @@ -2148,7 +2183,7 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in ] for wc in selected_wildcards: if self.__ppp.debug_level == DEBUG_LEVEL.full: - self.__ppp.logger.debug(f"Set default filter for wildcard '{wc}'") + 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) @@ -2172,7 +2207,7 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in filter_specifier: Optional[list[list[str]]] = None, wildcard_key: str = None, ) -> list[dict]: - msg_where = f"wildcard '{wildcard_key}'" if wildcard_key else "choices" + 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): @@ -2194,7 +2229,7 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in filtered_choice_values.append(c) if not filtered_choice_values: self.warn_or_stop( - f"Wildcard filter specifier '{','.join(['+'.join(y for y in x) for x in filter_specifier])}' found no matches in choices for wildcard '{wildcard_key}'!" + 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() @@ -2206,18 +2241,22 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in if cmd == "include": wcs = self.__ppp.wildcard_obj.get_wildcards(cmd_args) if not wcs: - self.warn_or_stop(f"Not found included wildcard '{cmd_args}' at {msg_where}!") + 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 '{wc.key}' at {msg_where} (chain starts at '{self.__seen_wildcards[0]}')!" + 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 '{wc.key}'") - self.__ppp.logger.debug(f"Including choices from wildcard '{wc.key}'") + 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) @@ -2229,7 +2268,7 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in } ) else: - self.warn_or_stop(f"Unsupported choice command '{cmd}' at {msg_where}!") + 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 @@ -2266,9 +2305,9 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in 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 '{wildcard_key}'" if wildcard_key else "choices" + msg_where = f"wildcard '{escape_single_quotes(wildcard_key)}'" if wildcard_key else "choices" if sampler != "~": - self.warn_or_stop(f"Unsupported sampler '{sampler}' at {msg_where} options!") + 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] = [] @@ -2317,7 +2356,7 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in 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 '{separator}'" if num_choices > 1 else "") + + (f" and separating with '{escape_single_quotes(separator)}'" if num_choices > 1 else "") ) ) if num_choices > 0: @@ -2364,7 +2403,7 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in suffix = "" results = [] if self.__ppp.debug_level == DEBUG_LEVEL.full: - list_unseen = [f"'{x}'" for x in self.__seen_wildcards[seen_wildcards_len:]] + 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) @@ -2473,7 +2512,7 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in ) except lark.exceptions.UnexpectedInput as e: self.warn_or_stop( - f"Error parsing choice prefix '{prefix}' in wildcard '{wildcard.key}'! : {e.__class__.__name__}", + 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) @@ -2484,7 +2523,7 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in ) except lark.exceptions.UnexpectedInput as e: self.warn_or_stop( - f"Error parsing choice suffix '{suffix}' in wildcard '{wildcard.key}'! : {e.__class__.__name__}", + f"Error parsing choice suffix '{escape_single_quotes(suffix)}' in wildcard '{escape_single_quotes(wildcard.key)}'! : {e.__class__.__name__}", e, ) n = 1 @@ -2517,7 +2556,7 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in ) except lark.exceptions.UnexpectedInput as e: self.warn_or_stop( - f"Error parsing condition '{condition}' in wildcard '{wildcard.key}'! : {e.__class__.__name__}", + f"Error parsing condition '{escape_single_quotes(condition)}' in wildcard '{escape_single_quotes(wildcard.key)}'! : {e.__class__.__name__}", e, ) cv["if"] = None @@ -2532,7 +2571,7 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in ) except lark.exceptions.UnexpectedInput as e: self.warn_or_stop( - f"Error parsing choice content '{content}' in wildcard '{wildcard.key}'! : {e.__class__.__name__}", + f"Error parsing choice content '{escape_single_quotes(content)}' in wildcard '{escape_single_quotes(wildcard.key)}'! : {e.__class__.__name__}", e, ) cv["content"] = None @@ -2541,9 +2580,13 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in self.__ppp.logger.debug(f"Processed choice {cv}") choice_values.append(cv) else: - self.warn_or_stop(f"Invalid choice {cv} in wildcard '{wildcard.key}'!") + 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 '{wildcard.key}'!") + self.warn_or_stop( + f"Invalid choice {cv} in wildcard '{escape_single_quotes(wildcard.key)}'!" + ) else: try: choice_values.append( @@ -2553,13 +2596,14 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in ) except lark.exceptions.UnexpectedInput as e: self.warn_or_stop( - f"Error parsing choice '{cv}' in wildcard '{wildcard.key}'! : {e.__class__.__name__}", e + 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 '{wildcard.key}' ({(t2-t1) / 1_000_000_000:.3f} seconds)" + f"Processed choices for wildcard '{escape_single_quotes(wildcard.key)}' ({(t2-t1) / 1_000_000_000:.3f} seconds)" ) return (options, choice_values) @@ -2618,7 +2662,7 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in and any(x.isdecimal() for x in filter_specifier) ): self.__ppp.logger.warning( - f"Using a globbing wildcard '{wildcard_key}' with positional index filters is not recommended!" + f"Using a globbing wildcard '{escape_single_quotes(wildcard_key)}' with positional index filters is not recommended!" ) var_object = tree.children[3] variablename = None @@ -2639,19 +2683,21 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in return if wildcard.key in self.__seen_wildcards: self.warn_or_stop( - f"Circular reference detected with wildcard '{self.__seen_wildcards[-1]}' (chain starts at '{self.__seen_wildcards[0]}')!" + 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 '{wildcard.key}'") + 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 '{wildcard.key}' are ignored!") + 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: @@ -2664,11 +2710,11 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in self.detectedWildcards.append(wc) self.result += wc if self.__ppp.debug_level == DEBUG_LEVEL.full: - list_unseen = [f"'{x}'" for x in self.__seen_wildcards[seen_wildcards_len:]] + 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"'{wc}'") + self.__debug_end("wildcard", start_result, t2 - t1, f"'{escape_single_quotes(wc)}'") def choices(self, tree: lark.Tree): """ @@ -2687,7 +2733,7 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in self.detectedWildcards.append(ch) self.result += ch t2 = time.monotonic_ns() - self.__debug_end("choices", start_result, t2 - t1, f"'{ch}'") + self.__debug_end("choices", start_result, t2 - t1, f"'{escape_single_quotes(ch)}'") def __default__(self, tree): t1 = time.monotonic_ns() diff --git a/ppp_cache.py b/ppp_cache.py index ec9d6ce..44ec45f 100644 --- a/ppp_cache.py +++ b/ppp_cache.py @@ -2,7 +2,7 @@ from collections import OrderedDict from logging import Logger from typing import Tuple -from ppp_logging import DEBUG_LEVEL # pylint: disable=import-error +from ppp_logging import DEBUG_LEVEL class PPPLRUCache: diff --git a/ppp_classes.py b/ppp_classes.py new file mode 100644 index 0000000..6922a34 --- /dev/null +++ b/ppp_classes.py @@ -0,0 +1,151 @@ +"""Pydantic models for the PPP configuration file structure (ppp_config.yaml).""" + +import re +from enum import Enum +from typing import Literal, Optional +from pydantic import BaseModel, ConfigDict, Field, field_validator, model_validator + +class SUPPORTED_APPS(Enum): + comfyui = "comfyui" + a1111 = "a1111" + forge = "forge" + reforge = "reforge" + sdnext = "sdnext" + tests = "tests" # for testing purposes only, not a real app + +SUPPORTED_APPS_NAMES = { + SUPPORTED_APPS.comfyui: "ComfyUI", + SUPPORTED_APPS.sdnext: "SD.Next", + SUPPORTED_APPS.forge: "Forge", + SUPPORTED_APPS.reforge: "reForge", + SUPPORTED_APPS.a1111: "A1111 (or compatible)", + SUPPORTED_APPS.tests: "Tests", +} + +# ------------------- Host configuration ------------------- + +AttentionOption = Literal["ok", "parentheses", "disable", "remove", "error"] +SchedulingOption = Literal["ok", "before", "after", "first", "remove", "error"] +AlternationOption = Literal["ok", "first", "remove", "error"] +AndOption = Literal["ok", "eol", "comma", "remove", "error"] +BreakOption = Literal["ok", "eol", "comma", "remove", "error"] + + +class HostConfig(BaseModel): + """Configuration for a specific host application.""" + + model_config = ConfigDict(populate_by_name=True, extra="forbid") + + attention: AttentionOption = "ok" + scheduling: SchedulingOption = "ok" + alternation: AlternationOption = "ok" + and_: AndOption = Field("ok", alias="and") + break_: BreakOption = Field("ok", alias="break") + + +# ------------------- Model detection ------------------- + + +class ModelDetectConfig(BaseModel): + """Detection configuration for a specific host when loading a model.""" + + model_config = ConfigDict(populate_by_name=True) + + class_: Optional[list[str]] = Field(None, alias="class") + property: Optional[str] = None + + @model_validator(mode="after") + def check_class_or_property(self) -> "ModelDetectConfig": + if self.class_ is None and self.property is None: + raise ValueError("either 'class' or 'property' must be specified") + return self + + +# ------------------- Variant find_in_filename ------------------- + + +class FindInFilenamePattern(BaseModel): + """A regex pattern with optional flags used to identify a model variant in the filename.""" + + regex: str + flags: int = 0 + + @field_validator("flags", mode="before") + @classmethod + def parse_flags(cls, v: object) -> int: + if isinstance(v, int): + return v + if isinstance(v, list): + flag_value = 0 + for flag in v: + if not isinstance(flag, str) or not hasattr(re, flag): + raise ValueError(f"invalid regex flag '{flag}'") + flag_value |= getattr(re, flag) + return flag_value + raise ValueError(f"expected int or list of flag-name strings, got {type(v).__name__}") + + @model_validator(mode="after") + def validate_regex(self) -> "FindInFilenamePattern": + try: + re.compile(self.regex, self.flags) + except re.error as exc: + raise ValueError(f"invalid regex pattern '{self.regex}': {exc}") from exc + return self + + +class VariantConfig(BaseModel): + """Configuration for a specific model variant.""" + + find_in_filename: list[FindInFilenamePattern] + + @field_validator("find_in_filename", mode="before") + @classmethod + def normalize_find_in_filename(cls, v: object) -> list: + """Normalize str / dict / list input to always be a list of FindInFilenamePattern-compatible dicts.""" + if isinstance(v, str): + return [{"regex": v, "flags": re.IGNORECASE}] + if isinstance(v, dict): + return [v] + if isinstance(v, list): + normalized = [] + for item in v: + if isinstance(item, str): + normalized.append({"regex": item, "flags": re.IGNORECASE}) + elif isinstance(item, dict): + normalized.append(item) + else: + raise ValueError(f"expected str or dict in 'find_in_filename' list, got {type(item).__name__}") + return normalized + raise ValueError(f"expected str, dict, or list for 'find_in_filename', got {type(v).__name__}") + + +# ------------------- Model configuration ------------------- + + +class ModelConfig(BaseModel): + """Configuration for a supported base model.""" + + detect: Optional[dict[str, Optional[ModelDetectConfig]]] = None + variants: Optional[dict[str, VariantConfig]] = None + + @model_validator(mode="after") + def check_detect_or_variants(self) -> "ModelConfig": + if self.detect is None and self.variants is None: + raise ValueError("at least one of 'detect' or 'variants' must be specified") + return self + + +# ------------------- Top-level configuration ------------------- + + +class PPPConfig(BaseModel): + """Top-level PPP configuration structure matching ppp_config.yaml.""" + + hosts: Optional[dict[str, Optional[HostConfig]]] = None + models: Optional[dict[str, Optional[ModelConfig | None]]] = None + + @model_validator(mode="after") + def check_hosts_or_models(self) -> "PPPConfig": + if self.hosts is None and self.models is None: + raise ValueError("at least one of 'hosts' or 'models' must be specified") + return self diff --git a/ppp_comfyui.py b/ppp_comfyui.py index a117aaa..e1b4447 100644 --- a/ppp_comfyui.py +++ b/ppp_comfyui.py @@ -1,14 +1,14 @@ import os -# pylint: disable=import-error -import folder_paths # type: ignore -import nodes # type: ignore +import folder_paths # pylint: disable=import-error # type: ignore +import nodes # pylint: disable=import-error # type: ignore -from .ppp import PromptPostProcessor -from .ppp_hosts import SUPPORTED_APPS -from .ppp_logging import DEBUG_LEVEL, PromptPostProcessorLogFactory -from .ppp_wildcards import PPPWildcards -from .ppp_enmappings import PPPExtraNetworkMappings +from ppp import PromptPostProcessor +from ppp_classes import SUPPORTED_APPS +from ppp_logging import DEBUG_LEVEL, PromptPostProcessorLogFactory +from ppp_utils import escape_single_quotes +from ppp_wildcards import PPPWildcards +from ppp_enmappings import PPPExtraNetworkMappings if __name__ == "__main__": raise SystemExit("This script must be run from ComfyUI") @@ -166,7 +166,7 @@ class PromptPostProcessorComfyUINode: for input_name, input_type in input_types.items(): t = expected[input_name] if input_type != t: - return f"Invalid type for input '{input_name}': {input_type} (expected {t})" + return f"Invalid type for input '{escape_single_quotes(input_name)}': {input_type} (expected {t})" return True RETURN_TYPES = ( diff --git a/ppp_enmappings.py b/ppp_enmappings.py index 3cf3321..c353f35 100644 --- a/ppp_enmappings.py +++ b/ppp_enmappings.py @@ -3,8 +3,8 @@ from typing import Optional import logging import yaml -from ppp_logging import DEBUG_LEVEL # pylint: disable=import-error -from ppp_utils import deep_freeze # pylint: disable=import-error +from ppp_logging import DEBUG_LEVEL +from ppp_utils import deep_freeze, escape_single_quotes class PPPENMappingVariant: @@ -209,25 +209,27 @@ 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 '{full_path}'!") + 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 '{kind}:*' in file '{full_path}'!") + 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 '{key}' in file '{full_path}'!" + 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 '{key}' in file '{full_path}' and '{self.extranetwork_mappings[key].file}'!" + 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 '{key}' in file '{full_path}'!" + 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) @@ -245,12 +247,16 @@ 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 '{full_path}' with utf-8 encoding, trying windows-1252...") + 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 '{full_path}': {e}") + 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): """ @@ -260,7 +266,9 @@ class PPPExtraNetworkMappings: directory (str): The path to the directory. """ if not os.path.exists(directory): - self.__logger.warning(f"Extra network mappings directory '{directory}' does not exist!") + 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_hosts.py b/ppp_hosts.py deleted file mode 100644 index 03ea527..0000000 --- a/ppp_hosts.py +++ /dev/null @@ -1,19 +0,0 @@ -from enum import Enum - - -class SUPPORTED_APPS(Enum): - comfyui = "comfyui" - a1111 = "a1111" - forge = "forge" - reforge = "reforge" - sdnext = "sdnext" - tests = "tests" # for testing purposes only, not a real app - -SUPPORTED_APPS_NAMES = { - SUPPORTED_APPS.comfyui: "ComfyUI", - SUPPORTED_APPS.sdnext: "SD.Next", - SUPPORTED_APPS.forge: "Forge", - SUPPORTED_APPS.reforge: "reForge", - SUPPORTED_APPS.a1111: "A1111 (or compatible)", - SUPPORTED_APPS.tests: "Tests", -} diff --git a/ppp_logging.py b/ppp_logging.py index 0696f9e..113f578 100644 --- a/ppp_logging.py +++ b/ppp_logging.py @@ -2,7 +2,7 @@ from enum import Enum import logging import sys import copy -from ppp_hosts import SUPPORTED_APPS # pylint: disable=import-error +from ppp_classes import SUPPORTED_APPS class DEBUG_LEVEL(Enum): diff --git a/ppp_utils.py b/ppp_utils.py index 0fd10b1..18ef055 100644 --- a/ppp_utils.py +++ b/ppp_utils.py @@ -15,3 +15,27 @@ def deep_freeze(obj): if isinstance(obj, set): return tuple(deep_freeze(i) for i in sorted(obj)) return obj + +def escape_single_quotes(s: str): + """ + Escape single quotes in a string. + + Args: + s (str): The string to escape. + + Returns: + str: The escaped string. + """ + return s.replace("'", "\\'") + +def escape_double_quotes(s: str): + """ + Escape double quotes in a string. + + Args: + s (str): The string to escape. + + Returns: + str: The escaped string. + """ + return s.replace('"', '\\"') diff --git a/ppp_wildcards.py b/ppp_wildcards.py index ec0959c..3125ad9 100644 --- a/ppp_wildcards.py +++ b/ppp_wildcards.py @@ -4,8 +4,8 @@ from typing import Optional import logging import yaml -from ppp_logging import DEBUG_LEVEL # pylint: disable=import-error -from ppp_utils import deep_freeze # pylint: disable=import-error +from ppp_logging import DEBUG_LEVEL +from ppp_utils import deep_freeze, escape_single_quotes class PPPWildcard: @@ -208,7 +208,7 @@ 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 '{full_path}': {e}") + self.__logger.error(f"Error reading wildcard file '{escape_single_quotes(full_path)}': {e}") def __get_wildcards_in_input(self, wildcards_input: str): """ @@ -284,7 +284,9 @@ class PPPWildcards: if isinstance(obj, (int, float, bool)): return [str(obj)] if not isinstance(obj, list) or len(obj) == 0: - self.__logger.warning(f"Invalid format in wildcard '{'/'.join(key_parts)}' in file '{full_path}'!") + self.__logger.warning( + f"Invalid format in wildcard '{escape_single_quotes('/'.join(key_parts))}' in file '{escape_single_quotes(full_path)}'!" + ) return None choices = [] for i, c in enumerate(obj): @@ -297,7 +299,7 @@ class PPPWildcards: choices.append(self.__process_dict_choice(c, full_path, key_parts, i)) else: self.__logger.warning( - f"Invalid choice {i+1} in wildcard '{'/'.join(key_parts)}' in file '{full_path}'!" + f"Invalid choice {i+1} in wildcard '{escape_single_quotes('/'.join(key_parts))}' in file '{escape_single_quotes(full_path)}'!" ) return choices @@ -328,7 +330,9 @@ 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 '{'/'.join(key_parts)}' in file '{full_path}'!") + 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): @@ -371,16 +375,18 @@ class PPPWildcards: fullkey = "/".join(tmp_key_parts) if self.wildcards.get(fullkey, None) is not None: self.__logger.warning( - f"Duplicate wildcard '{fullkey}' in file '{full_path}' and '{self.wildcards[fullkey].file}'!" + 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 '{fullkey}' in file '{full_path}'!") + 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 '{fullkey}' in file '{full_path}'! (cannot start with underscore)" + 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) @@ -390,20 +396,22 @@ class PPPWildcards: elif isinstance(content, (int, float, bool)): content = [str(content)] if not isinstance(content, list): - self.__logger.warning(f"Invalid wildcard in file '{full_path}'!") + 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 '{fullkey}' in file '{full_path}' and '{self.wildcards[fullkey].file}'!" + 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 '{fullkey}' in file '{full_path}'!") + 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 '{fullkey}' in file '{full_path}'! (cannot start with underscore)" + 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) @@ -422,7 +430,9 @@ 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 '{full_path}' with utf-8 encoding, trying windows-1252...") + 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) @@ -441,7 +451,9 @@ 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 '{full_path}' with utf-8 encoding, trying windows-1252...") + 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)) @@ -457,7 +469,7 @@ class PPPWildcards: directory (str): The path to the directory. """ if not os.path.exists(directory): - self.__logger.warning(f"Wildcard directory '{directory}' does not exist!") + 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 5d8f478..5868102 100644 --- a/scripts/ppp_script.py +++ b/scripts/ppp_script.py @@ -14,12 +14,12 @@ from modules.processing import StableDiffusionProcessing # pylint: disable=impo 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 # pylint: disable=import-error -from ppp_hosts import SUPPORTED_APPS, SUPPORTED_APPS_NAMES # pylint: disable=import-error -from ppp_logging import DEBUG_LEVEL, PromptPostProcessorLogFactory # pylint: disable=import-error -from ppp_cache import PPPLRUCache # pylint: disable=import-error -from ppp_wildcards import PPPWildcards # pylint: disable=import-error -from ppp_enmappings import PPPExtraNetworkMappings # pylint: disable=import-error +from ppp import PromptPostProcessor +from ppp_classes import SUPPORTED_APPS, SUPPORTED_APPS_NAMES +from ppp_logging import DEBUG_LEVEL, PromptPostProcessorLogFactory +from ppp_cache import PPPLRUCache +from ppp_wildcards import PPPWildcards +from ppp_enmappings import PPPExtraNetworkMappings class PromptPostProcessorA1111Script(scripts.Script): @@ -294,7 +294,7 @@ class PromptPostProcessorA1111Script(scripts.Script): self.wildcards_obj, self.extranetwork_mappings_obj, ) - hash_options = ppp.options_hash() + hash_options = ppp.options_hash() hash_envinfo = ppp.envinfo_hash() prompts_list = [] diff --git a/tests/base_tests.py b/tests/base_tests.py new file mode 100644 index 0000000..b62296a --- /dev/null +++ b/tests/base_tests.py @@ -0,0 +1,207 @@ +import os +import logging +from typing import NamedTuple, Optional +import unittest +import datetime + +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 +from ppp_logging import DEBUG_LEVEL, PromptPostProcessorLogFactory # pylint: disable=import-error + + +class PromptPair(NamedTuple): + prompt: str = "" + negative_prompt: str = "" + + +class TestPromptPostProcessorBase(unittest.TestCase): + """ + A test case class for testing the PromptPostProcessor class. + """ + + def setUp(self, enable_file_logging=False): + """ + Set up the test case by initializing the necessary objects and configurations. + + Args: + enable_file_logging (bool): Whether to enable logging to a file. Defaults to True. + """ + self.enable_file_logging = enable_file_logging + test_name = self.id().split(".")[-1] # Extract the test method name + timestamp = datetime.datetime.now().strftime("%Y%m%d_%H%M%S") + + if self.enable_file_logging: + log_filename = f"tests/logs/{test_name}_{timestamp}.log" + else: + log_filename = None # Disable file logging + + self.lf = PromptPostProcessorLogFactory(None, 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.def_env_info = { + "app": "tests", + "ppp_config": None, + "model_class": "SDXL", + "property_base": {"is_sdxl": True}, + "models_path": "./webui/models", + "model_filename": "./webui/models/Stable-diffusion/testmodel.safetensors", + } + self.interrupted = False + self.wildcards_obj = PPPWildcards(self.lf.log) + self.extranetwork_maps_obj = PPPExtraNetworkMappings(self.lf.log) + self.wildcards_obj.refresh_wildcards( + DEBUG_LEVEL.full, + [ + os.path.abspath(os.path.join(os.path.dirname(__file__), "wildcards")), + os.path.abspath(os.path.join(os.path.dirname(__file__), "wildcards2")), + ], + """ + yaml_input: + wildcardI: + - choice1 + - choice2 + - choice3 + """, + ) + self.extranetwork_maps_obj.refresh_extranetwork_mappings( + DEBUG_LEVEL.full, + [os.path.abspath(os.path.join(os.path.dirname(__file__), "enmappings"))], + """ + """, + ) + grammar_filename = os.path.join(os.path.dirname(os.path.realpath(__file__)), "../grammar.lark") + with open(grammar_filename, "r", encoding="utf-8") as file: + self.grammar_content = file.read() + + def interrupt(self): + self.interrupted = True + + def process( + self, + input_prompts: PromptPair, + expected_output_prompts: Optional[PromptPair | list[PromptPair]] = None, + seed: int = 1, + ppp: Optional[str | PromptPostProcessor] = None, + interrupted: bool = False, + variables: dict[str, str] | None = None, + ): + """ + Process the prompt and compare the results with the expected prompts. + + Args: + input_prompts (PromptPair): The input prompts. + expected_output_prompts (PromptPair | list[PromptPair], optional): The expected prompts. + seed (int, optional): The seed value. Defaults to 1. + ppp (Optional[str | PromptPostProcessor], optional): The PromptPostProcessor instance or type. Defaults to None. + interrupted (bool, optional): The interrupted flag. Defaults to False. + variables (dict[str,str]|None, optional): Output variables to check. Defaults to None. + + Returns: + None + """ + if isinstance(ppp, str): + 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, + }, + self.grammar_content, + 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.wildcards_obj, + self.extranetwork_maps_obj, + ) + out = ( + [PromptPair("", "")] + if expected_output_prompts is None + else expected_output_prompts if isinstance(expected_output_prompts, list) else [expected_output_prompts] + ) + for eo in out: + result_prompt, result_negative_prompt, output_variables = the_obj.process_prompt( + input_prompts.prompt, + input_prompts.negative_prompt, + seed, + ) + self.assertEqual(self.interrupted, interrupted, "Interrupted flag is incorrect") + if not self.interrupted: + if expected_output_prompts is not None: + self.assertEqual(result_prompt, eo.prompt, "Incorrect prompt") + self.assertEqual(result_negative_prompt, eo.negative_prompt, "Incorrect negative prompt") + if variables is not None: + for var_name, var_value in variables.items(): + self.assertIn( + var_name, output_variables, f"Variable '{var_name}' not found in output variables" + ) + self.assertEqual( + output_variables[var_name], + var_value, + f"Variable '{var_name}' has incorrect value", + ) + seed += 1 diff --git a/tests/tests.py b/tests/tests.py deleted file mode 100644 index 8a5cc34..0000000 --- a/tests/tests.py +++ /dev/null @@ -1,1961 +0,0 @@ -import os -import logging -from typing import NamedTuple, Optional -import unittest -import datetime - -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 -from ppp_logging import DEBUG_LEVEL, PromptPostProcessorLogFactory # pylint: disable=import-error - - -class PromptPair(NamedTuple): - prompt: str = "" - negative_prompt: str = "" - - -class TestPromptPostProcessorBase(unittest.TestCase): - """ - A test case class for testing the PromptPostProcessor class. - """ - - def setUp(self, enable_file_logging=False): - """ - Set up the test case by initializing the necessary objects and configurations. - - Args: - enable_file_logging (bool): Whether to enable logging to a file. Defaults to True. - """ - self.enable_file_logging = enable_file_logging - test_name = self.id().split(".")[-1] # Extract the test method name - timestamp = datetime.datetime.now().strftime("%Y%m%d_%H%M%S") - - if self.enable_file_logging: - log_filename = f"tests/logs/{test_name}_{timestamp}.log" - else: - log_filename = None # Disable file logging - - self.lf = PromptPostProcessorLogFactory(None, 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.def_env_info = { - "app": "tests", - "ppp_config": None, - "model_class": "SDXL", - "property_base": {"is_sdxl": True}, - "models_path": "./webui/models", - "model_filename": "./webui/models/Stable-diffusion/testmodel.safetensors", - } - self.interrupted = False - self.wildcards_obj = PPPWildcards(self.lf.log) - self.extranetwork_maps_obj = PPPExtraNetworkMappings(self.lf.log) - self.wildcards_obj.refresh_wildcards( - DEBUG_LEVEL.full, - [ - os.path.abspath(os.path.join(os.path.dirname(__file__), "wildcards")), - os.path.abspath(os.path.join(os.path.dirname(__file__), "wildcards2")), - ], - """ - yaml_input: - wildcardI: - - choice1 - - choice2 - - choice3 - """, - ) - self.extranetwork_maps_obj.refresh_extranetwork_mappings( - DEBUG_LEVEL.full, - [os.path.abspath(os.path.join(os.path.dirname(__file__), "enmappings"))], - """ - """, - ) - grammar_filename = os.path.join(os.path.dirname(os.path.realpath(__file__)), "../grammar.lark") - with open(grammar_filename, "r", encoding="utf-8") as file: - self.grammar_content = file.read() - - def interrupt(self): - self.interrupted = True - - def process( - self, - input_prompts: PromptPair, - expected_output_prompts: Optional[PromptPair | list[PromptPair]] = None, - seed: int = 1, - ppp: Optional[str | PromptPostProcessor] = None, - interrupted: bool = False, - variables: dict[str, str] | None = None, - ): - """ - Process the prompt and compare the results with the expected prompts. - - Args: - input_prompts (PromptPair): The input prompts. - expected_output_prompts (PromptPair | list[PromptPair], optional): The expected prompts. - seed (int, optional): The seed value. Defaults to 1. - ppp (Optional[str | PromptPostProcessor], optional): The PromptPostProcessor instance or type. Defaults to None. - interrupted (bool, optional): The interrupted flag. Defaults to False. - variables (dict[str,str]|None, optional): Output variables to check. Defaults to None. - - Returns: - None - """ - if isinstance(ppp, str): - 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, - }, - self.grammar_content, - 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.wildcards_obj, - self.extranetwork_maps_obj, - ) - out = ( - [PromptPair("", "")] - if expected_output_prompts is None - else expected_output_prompts if isinstance(expected_output_prompts, list) else [expected_output_prompts] - ) - for eo in out: - result_prompt, result_negative_prompt, output_variables = the_obj.process_prompt( - input_prompts.prompt, - input_prompts.negative_prompt, - seed, - ) - self.assertEqual(self.interrupted, interrupted, "Interrupted flag is incorrect") - if not self.interrupted: - if expected_output_prompts is not None: - self.assertEqual(result_prompt, eo.prompt, "Incorrect prompt") - self.assertEqual(result_negative_prompt, eo.negative_prompt, "Incorrect negative prompt") - if variables is not None: - for var_name, var_value in variables.items(): - self.assertIn( - var_name, output_variables, f"Variable '{var_name}' not found in output variables" - ) - self.assertEqual( - output_variables[var_name], - var_value, - f"Variable '{var_name}' has incorrect value", - ) - seed += 1 - - -class TestPromptPostProcessor(TestPromptPostProcessorBase): - - def setUp(self): # pylint: disable=arguments-differ - super().setUp(enable_file_logging=False) - - # Send To Negative tests - - def test_stn_simple(self): # negtags with different parameters and separations - self.process( - PromptPair( - "flowersred, green, blueyellow, purpleblack", - "normal quality, worse quality", - ), - PromptPair("flowers", "red, green, yellow, normal quality, purple, worse quality, black, blue"), - ) - - def test_stn_complex(self): # complex negtags - self.process( - PromptPair( - "red ((pink)), flowers purple, mauveblue, yellow green", - "normal quality, , bad quality, worse quality", - ), - PromptPair( - "flowers", - "red, (pink:1.21), normal quality, mauve, yellow, bad quality, green, worse quality, purple, blue", - ), - ) - - def test_stn_complex_nocleanup(self): # complex negtags with no cleanup - self.process( - PromptPair( - "red ((pink)), flowers purple, mauveblue, yellow green", - "normal quality, , bad quality, worse quality", - ), - PromptPair( - " (()), flowers , , ", - "red, ((pink)), normal quality, mauve, yellow, bad quality, green, worse quality, purple, blue", - ), - ppp="nocup", - ) - - def test_stn_inside_attention(self): # negtag inside attention - self.process( - PromptPair( - "[neg1] this is a ((testneg2) (test:2.0): 1.5 ) (red[square]:1.5)", - "normal quality", - ), - PromptPair( - "this is a ((test) (test:2):1.5) (red:1.5)", "[neg1], ([square]:1.5), normal quality, (neg2:1.65)" - ), - ) - - def test_stn_inside_alternation(self): # negtag inside alternation - self.process( - PromptPair( - "this is a (([complexneg1|simpleneg2|regularneg3] test)(test:2.0):1.5)", - "normal quality", - ), - PromptPair( - "this is a (([complex|simple|regular] test)(test:2):1.5)", - "([neg1||]:1.65), ([|neg2|]:1.65), ([||neg3]:1.65), normal quality", - ), - ) - - def test_stn_inside_alternation_recursive(self): # negtag inside alternation (recursive alternation) - self.process( - PromptPair( - "this is a (([complexneg1[one|twoneg12||three|four(neg14)]|simpleneg2|regularneg3] test)(test:2.0):1.5)", - "normal quality", - ), - PromptPair( - "this is a (([complex[one|two||three|four]|simple|regular] test)(test:2):1.5)", - "([neg1||]:1.65), ([[|neg12|||]||]:1.65), ([[||||(neg14)]||]:1.65), ([|neg2|]:1.65), ([||neg3]:1.65), normal quality", - ), - ) - - def test_stn_inside_scheduling(self): # negtag inside scheduling - self.process( - PromptPair("this is [abcneg1:defneg2: 5 ]", "normal quality"), - [PromptPair("this is [abc:def:5]", "[neg1::5], normal quality, [neg2:5]")], - ) - - def test_stn_complex_features(self): # complex negtags with AND, BREAK and other features - self.process( - PromptPair( - "[neg5] this \\(is\\): a (([complex|simpleneg6|regular] testneg1)(test:2.0):1.5) \nBREAK, BREAK with [abcneg4:defneg2(neg3:1.6):5]:0.5 AND loratrigger AND AND hypernettrigger :0.3", - "normal quality, ", - ), - PromptPair( - "this \\(is\\): a (([complex|simple|regular] test)(test:2):1.5)\nBREAK with [abc:def:5]:0.5 AND loratrigger AND hypernettrigger :0.3", - "[neg5], ([|neg6|]:1.65), (neg1:1.65), [neg4::5], normal quality, [neg2(neg3:1.6):5]", - ), - ) - - def test_stn_complex_features_newformat(self): # complex negtags with AND, BREAK and other features (new format) - self.process( - PromptPair( - "[neg5] this \\(is\\): a (([complex|simpleneg6|regular] testneg1)(test:2.0):1.5) \nBREAK, BREAK with [abcneg4:defneg2(neg3:1.6):5]:0.5 AND loratrigger AND AND hypernettrigger :0.3", - "normal quality, ", - ), - PromptPair( - "this \\(is\\): a (([complex|simple|regular] test)(test:2):1.5)\nBREAK with [abc:def:5]:0.5 AND loratrigger AND hypernettrigger :0.3", - "[neg5], ([|neg6|]:1.65), (neg1:1.65), [neg4::5], normal quality, [neg2(neg3:1.6):5]", - ), - ) - - def test_stn_inside_alternation_recursive_2(self): # negtag inside alternation (recursive alternation) - self.process( - PromptPair( - "[pos1neg1[pos11|pos12neg12||pos14|pos15neg15]|pos2neg2|pos3neg3]", - "", - ), - PromptPair( - "[pos1[pos11|pos12||pos14|pos15]|pos2|pos3]", - "[neg1||], [[|neg12|||]||], [[||||neg15]||], [|neg2|], [||neg3]", - # "[neg1[|neg12|||neg15]|neg2|neg3]", # expected output if the constructs were unified - ), - ) - - # Cleanup tests - - def test_cl_simple(self): # simple cleanup - self.process( - PromptPair(" this is a ((test ), , , (), , [] ( , test ,:2.0):1.5), (red:1.5) ", " normal quality "), - PromptPair("this is a ((test), (test,:2):1.5), (red:1.5)", "normal quality"), - ) - - def test_cl_complex(self): # complex cleanup - self.process( - PromptPair( - " this is BREAKABLE a ((test)), ,AND AND(() [] ANDERSON (test:2.0):1.5) :o BREAK \n BREAK (red:1.5) ", - " [:hands, feet, :0.15]normal quality ", - ), - PromptPair( - "this is BREAKABLE a (test:1.21) AND( ANDERSON (test:2):1.5) :o BREAK (red:1.5)", - "[:hands, feet, :0.15]normal quality", - ), - ) - - def test_cl_removenetworktags(self): # remove network tags - self.process( - PromptPair("this is a test__yaml/wildcard7__", ""), - PromptPair("this is a test", ""), - ppp=PromptPostProcessor( - self.ppp_logger, - self.interrupt, - self.def_env_info, - {**self.defopts, "remove_extranetwork_tags": True}, - self.grammar_content, - self.wildcards_obj, - self.extranetwork_maps_obj, - ), - ) - - def test_cl_dontremoveseparatorsoneol(self): # don't remove separators on eol - self.process( - PromptPair("this is a test,\nsecond line", ""), - 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, - }, - self.grammar_content, - self.wildcards_obj, - self.extranetwork_maps_obj, - ), - ) - - def test_cl_separatorswitheol(self): # don't remove eols with the separators - self.process( - PromptPair( - """{ (d:0.9) ,, (l:1.1) | (l:1.1) (d:0.9),,, } - (l:1.1) - (d:0.9)""", - "", - ), - PromptPair( - """ (l:1.1) (d:0.9), - (l:1.1) - (d:0.9)""", - "", - ), - 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, - }, - self.grammar_content, - self.wildcards_obj, - self.extranetwork_maps_obj, - ), - ) - - def test_cl_mergeattention(self): # merge attention - self.process( - PromptPair( - "this is (a test:0.9) of (attention (merging:1.2)) where ((this)) ((is joined:1.2)) and ([this too]:1.3)", - "", - ), - PromptPair( - "this is [a test] of (attention (merging:1.2)) where (this:1.21) (is joined:1.32) and (this too:1.17)", - "", - ), - ) - - def test_cl_not_mergeattention(self): # not merge attention - self.process( - PromptPair( - "this is (a test:0.9) of not (attention (merging:1.2)) where ((this)) ((is not joined:1.2)) and neither is ([this]:1.3)", - "", - ), - PromptPair( - "this is (a test:0.9) of not (attention (merging:1.2)) where ((this)) ((is not joined:1.2)) and neither is ([this]:1.3)", - "", - ), - ppp="nocup", - ) - - # Command tests - - def test_cmd_stn_complex_features(self): # complex stn command with AND, BREAK and other features - self.process( - PromptPair( - "[neg5] this \\(is\\): a (([complex|simpleneg6|regular] testneg1)(test:2.0):1.5) \nBREAK, BREAK with [abcneg4:defneg2(neg3:1.6):5]:0.5 AND loratrigger AND AND hypernettrigger :0.3", - "normal quality, ", - ), - PromptPair( - "this \\(is\\): a (([complex|simple|regular] test)(test:2):1.5)\nBREAK with [abc:def:5]:0.5 AND loratrigger AND hypernettrigger :0.3", - "[neg5], ([|neg6|]:1.65), (neg1:1.65), [neg4::5], normal quality, [neg2(neg3:1.6):5]", - ), - ) - - def test_cmd_if_complex_features(self): # complex if command - self.process( - PromptPair( - "this \\(is\\): a (([complex|simple|regular] test)(test:2.0):1.5) \nBREAK, BREAK with [abcneg4:def:5]:0.5 AND loratrigger hypernettrigger nothing:0.3", - "normal quality", - ), - PromptPair( - "this \\(is\\): a (([complex|simple|regular] test)(test:2):1.5)\nBREAK :0.5 AND hypernettrigger :0.3", - "normal quality", - ), - ) - - def test_cmd_if_nested(self): # nested if command - self.process( - PromptPair( - "this is SD1PONYSD2NOPONYNOPONY", - "", - ), - 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.wildcards_obj, - self.extranetwork_maps_obj, - ), - ) - - def test_cmd_set_if(self): # set and if commands - self.process( - PromptPair("valuethis test is OKnot OK", ""), - PromptPair("this test is OK", ""), - ) - - def test_cmd_set_empty(self): # set to empty - self.process( - PromptPair("${v2=}this test is not OKOK", ""), - PromptPair("this test is OK", ""), - ) - - def test_cmd_set_eval_if(self): # set and if commands - self.process( - PromptPair("valuethis test is OKnot OK", ""), - PromptPair("this test is OK", ""), - ) - - def test_cmd_set_if_echo_nested(self): # nested set, if and echo commands - self.process( - PromptPair( - "1OKnot OK NOK OK", - "", - ), - PromptPair("OK OK OK", ""), - ) - - def test_cmd_set_if_complex_conditions_1(self): # complex conditions (or) - self.process( - PromptPair( - "truefalsethis test is OKnot OK", - "", - ), - PromptPair("this test is OK", ""), - ) - - def test_cmd_set_if_complex_conditions_2(self): # complex conditions (and) - self.process( - PromptPair( - "truetruethis test is OKnot OK", - "", - ), - PromptPair("this test is OK", ""), - ) - - def test_cmd_set_if_complex_conditions_3(self): # complex conditions (not) - self.process( - PromptPair("falsethis test is OKnot OK", ""), - PromptPair("this test is OK", ""), - ) - - def test_cmd_set_if_complex_conditions_4(self): # complex conditions (not, precedence) - self.process( - PromptPair( - "truefalsethis test is OKnot OK", - "", - ), - PromptPair("this test is OK", ""), - ) - - def test_cmd_set_if_complex_conditions_5(self): # complex conditions (not, precedence, comparison) - self.process( - PromptPair( - "1falsethis test is OKnot OK", - "", - ), - PromptPair("this test is OK", ""), - ) - - def test_cmd_set_if_complex_conditions_6(self): # complex conditions - self.process( - PromptPair( - "123this test is OKnot OK", - "", - ), - PromptPair("this test is OK", ""), - ) - - def test_cmd_set_if_complex_conditions_7(self): # complex conditions - self.process( - PromptPair( - "123this test is OKnot OK", - "", - ), - PromptPair("this test is OK", ""), - ) - - def test_var_nested_1(self): # variable default nested in variable set - self.process( - PromptPair( - "${v1=test ${v2:OK}}${v1}", - "", - ), - PromptPair("test OK", ""), - variables={"v1": "test OK", "v2": "OK"}, - ) - - def test_var_nested_2(self): # variable set nested in variable default - self.process( - PromptPair( - "${v1:test ${v2=OK}${v2}}", - "", - ), - PromptPair("test OK", ""), - variables={"v1": "test OK", "v2": "OK"}, - ) - - def test_var_nested_3(self): # variable default nested in variable default - self.process( - PromptPair( - "${v1:test ${v2:OK}}", - "", - ), - PromptPair("test OK", ""), - variables={"v1": "test OK", "v2": "OK"}, - ) - - def test_cmd_set_if2(self): # set and more complex if commands - self.process( - PromptPair( - "First: value1this test is OKOK2not OK\nSecond: value3this test is OKnot OK", - "", - ), - PromptPair("First: this test is OK\nSecond: this test is OK", ""), - ) - - def test_cmd_set_add_if(self): # set, add and if commands - self.process( - PromptPair( - "value2this test is OKnot OK", - "", - ), - PromptPair("this test is OK", ""), - ) - - def test_cmd_set_add_DP_if(self): # set, add (DP format) and if commands - self.process( - PromptPair( - "${v=value}${v+=2}this test is OKnot OK", - "", - ), - PromptPair("this test is OK", ""), - ) - - def test_cmd_set_immediateeval(self): # set (DP format) with mixed evaluation - self.process( - PromptPair( - "${var=!__yaml/wildcard1__}the choices are: ${var}, ${var}, ${var2:default}, ${var3=__yaml/wildcard1__}${var3}, ${var3}", - "", - ), - PromptPair("the choices are: choice2, choice2, default, choice3, choice1", ""), - ppp="nocup", - ) - - def test_cmd_set_mixeval(self): # set and add (DP format) with mixed evaluation - self.process( - PromptPair( - "${var=__yaml/wildcard1__}the choices are: ${var}, ${var}, ${var+=, __yaml/wildcard2__}${var}, ${var}, ${var+=!, __yaml/wildcard3__}${var}, ${var}", - "", - ), - PromptPair( - "the choices are: choice2, choice3, choice1, choice1- choice2 -choice3, choice2, choice2 -choice1-choice3, choice2, choice3-choice1- choice2 , choice1, choice2 , choice2, choice3-choice1- choice2 , choice1, choice2 ", - "", - ), - ppp="nocup", - ) - - def test_cmd_set_ifundefined_if(self): # set, ifundefined and if commands - self.process( - PromptPair( - "valuethis test is OKnot OK", - "", - ), - PromptPair("this test is OK", ""), - ) - - def test_cmd_set_ifundefined_if_2(self): # set, ifundefined and if commands - self.process( - PromptPair( - "valuevalue2this test is OKnot OK", - "", - ), - PromptPair("this test is OK", ""), - ) - - def test_cmd_set_ifundefined_DP_if(self): # set, ifundefined (DP format) and if commands - self.process( - PromptPair( - "${v?=value}this test is OKnot OK", - "", - ), - PromptPair("this test is OK", ""), - ) - - def test_cmd_set_ifundefined_DP_if_2(self): # set, ifundefined (DP format) and if commands - self.process( - PromptPair( - "${v=!value}${v?=!value2}this test is OKnot OK", - "", - ), - PromptPair("this test is OK", ""), - ) - - def test_cmd_ext(self): # ext - self.process( - PromptPair( - "trigger1trigger2trigger4trigger5", - "", - ), - PromptPair( - "trigger1,trigger2,trigger4,trigger5", - "", - ), - ) - - def test_cmd_ext_map_notrigger(self): # ext mapping, no trigger - self.process( - PromptPair( - "", - "", - ), - PromptPair("triggergeneric1, triggergeneric2, two, triggergeneric1, triggergeneric2, two", ""), - ) - - def test_cmd_ext_map1(self): # ext mapping, no lora - self.process( - PromptPair( - "inlinetrigger", - "", - ), - PromptPair("inlinetrigger, triggergeneric1, triggergeneric2, two", ""), - ) - - def test_cmd_ext_map2(self): # ext mapping, lora with weight - self.process( - PromptPair( - "inlinetrigger", - "", - ), - 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.wildcards_obj, - self.extranetwork_maps_obj, - ), - ) - - def test_cmd_ext_map3(self): # ext mapping, lora with weight adjusted - self.process( - PromptPair( - "inlinetrigger", - "", - ), - 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.wildcards_obj, - self.extranetwork_maps_obj, - ), - ) - - def test_cmd_ext_map4(self): # ext mapping, lora with parameters - self.process( - PromptPair( - "inlinetrigger", - "", - ), - 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.wildcards_obj, - self.extranetwork_maps_obj, - ), - ) - - def test_cmd_ext_map5(self): # ext mapping, lora with no parameters - self.process( - PromptPair( - "inlinetrigger", - "", - ), - 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.wildcards_obj, - self.extranetwork_maps_obj, - ), - ) - - # Choices tests - - def test_ch_choices(self): # simple choices with weights - self.process( - PromptPair("the choices are: {3::choice1|2::choice2|choice3}", ""), - PromptPair("the choices are: choice2", ""), - ppp="nocup", - ) - - def test_ch_unsupportedsampler(self): # unsupported sampler - self.process( - PromptPair("the choices are: {@choice1|choice2|choice3}", ""), - PromptPair("", ""), - ppp="nocup", - interrupted=True, - ) - - def test_ch_choices_withcomments(self): # choices with comments and multiline - self.process( - PromptPair( - "the choices are: {\n3::choice1 # this is option 1\n|2::choice2\n# this was option 2\n|choice3 # this is option 3\n}", - "", - ), - PromptPair("the choices are: choice2", ""), - ppp="nocup", - ) - - def test_ch_choices_multiple(self): # choices with multiple selection - self.process( - PromptPair("the choices are: {~2$$, $$3::choice1|2:: choice2 |choice3}", ""), - PromptPair("the choices are: choice2 , choice3", ""), - ppp="nocup", - ) - - def test_ch_choices_if_multiple(self): # choices with if and multiple selection - self.process( - PromptPair("the choices are: {2$$, $$3::choice1|2 if _is_sd1::choice2|choice3}", ""), - PromptPair("the choices are: choice1, choice3", ""), - ppp="nocup", - ) - - def test_ch_choices_set_if_multiple(self): # choices with if user variable and multiple selection - self.process( - PromptPair("${var=test}the choices are: {2$$, $$3::choice1|2 if not var eq 'test'::choice2|choice3}", ""), - PromptPair("the choices are: choice1, choice3", ""), - ppp="nocup", - ) - - def test_ch_choices_set_if_nested(self): # nested choices with if user variable and multiple selection - self.process( - PromptPair( - "${var=test}the choices are: {2$$, $$3::choice1${var2=test2} {if var2 eq 'test2'::choice11|choice12}|2 if not var eq 'test'::choice2|choice3}", - "", - ), - PromptPair("the choices are: choice1 choice11, choice3", ""), - ppp="nocup", - ) - - def test_ch_choicesinsidelora(self): # simple choices inside a lora - self.process( - PromptPair("", ""), - PromptPair("", ""), - ppp="nocup", - ) - - def test_ch_removelorawithchoices(self): - self.process( - PromptPair("", ""), - PromptPair("", ""), - ppp=PromptPostProcessor( - self.ppp_logger, - self.interrupt, - self.def_env_info, - {**self.defopts, "remove_extranetwork_tags": True}, - self.grammar_content, - self.wildcards_obj, - self.extranetwork_maps_obj, - ), - ) - - def test_ch_cmd_includewildcard(self): - self.process( - PromptPair("{ch_one|ch_two|%0.5::include yaml/wildcard1}", ""), - PromptPair("ch_two", ""), - ppp="nocup", - ) - - # Wildcards tests - - def test_wc_ignore(self): # wildcards with ignore option - self.process( - PromptPair("__bad_wildcard__", "{option1|option2}"), - 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, - }, - self.grammar_content, - self.wildcards_obj, - self.extranetwork_maps_obj, - ), - ) - - def test_wc_remove(self): # wildcards with remove option - self.process( - PromptPair( - "[neg5] this is: __bad_wildcard__ a (([complex|simpleneg6|regular] testneg1)(test:2.0):1.5) \nBREAK, BREAK with [abcneg4:defneg2(neg3:1.6):5] ", - "normal quality, {option1|option2}", - ), - PromptPair( - "this is: a (([complex|simple|regular] test)(test:2):1.5)\nBREAK with [abc:def:5]", - "[neg5], ([|neg6|]:1.65), (neg1:1.65), [neg4::5], normal quality, [neg2(neg3:1.6):5]", - ), - ppp=PromptPostProcessor( - self.ppp_logger, - self.interrupt, - self.def_env_info, - { - **self.defopts, - "process_wildcards": False, - "if_wildcards": PromptPostProcessor.IFWILDCARDS_CHOICES.remove.value, - }, - self.grammar_content, - self.wildcards_obj, - self.extranetwork_maps_obj, - ), - ) - - def test_wc_warn(self): # wildcards with warn option - self.process( - PromptPair("__bad_wildcard__", "{option1|option2}"), - 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, - }, - self.grammar_content, - self.wildcards_obj, - self.extranetwork_maps_obj, - ), - ) - - def test_wc_stop(self): # wildcards with stop option - self.process( - PromptPair("__bad_wildcard__", "{option1|option2}"), - PromptPair( - PromptPostProcessor.WILDCARD_STOP.format("__bad_wildcard__") + "__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.stop.value, - }, - self.grammar_content, - self.wildcards_obj, - self.extranetwork_maps_obj, - ), - interrupted=True, - ) - - def test_wcinvar_warn(self): # wildcards in var with warn option - self.process( - PromptPair("${v=__bad_wildcard__}${v}", ""), - 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, - }, - self.grammar_content, - self.wildcards_obj, - self.extranetwork_maps_obj, - ), - ) - - def test_wc_invalid_name(self): - self.process( - PromptPair("the choices are: ___invalid__", ""), - PromptPair("the choices are: ___invalid__", ""), - ppp="nocup", - ) - - def test_wc_wildcard1a_text(self): # simple text wildcard - self.process( - PromptPair("the choices are: __text/wildcard1__", ""), - PromptPair("the choices are: choice2", ""), - ppp="nocup", - ) - - def test_wc_wildcard1a_json(self): # simple json wildcard - self.process( - PromptPair("the choices are: __json/wildcard1__", ""), - PromptPair("the choices are: choice2", ""), - ppp="nocup", - ) - - def test_wc_wildcard1a_yaml(self): # simple yaml wildcard - self.process( - PromptPair("the choices are: __yaml/wildcard1__", ""), - PromptPair("the choices are: choice2", ""), - ppp="nocup", - ) - - def test_wc_wildcard1b_text(self): # simple text wildcard with multiple choices - self.process( - PromptPair("the choices are: __2-$$text/wildcard1__", ""), - PromptPair("the choices are: choice3, choice1", ""), - ppp="nocup", - ) - - def test_wc_wildcard1b_json(self): # simple json wildcard with multiple choices - self.process( - PromptPair("the choices are: __2-$$json/wildcard1__", ""), - PromptPair("the choices are: choice3, choice1", ""), - ppp="nocup", - ) - - def test_wc_wildcard1b_yaml(self): # simple yaml wildcard with multiple choices - self.process( - PromptPair("the choices are: __2-$$yaml/wildcard1__", ""), - PromptPair("the choices are: choice3, choice1", ""), - ppp="nocup", - ) - - def test_wc_wildcard2_text(self): # simple text wildcard with default options - self.process( - PromptPair("the choices are: __text/wildcard2__", ""), - PromptPair("the choices are: choice3-choice1", ""), - ppp="nocup", - ) - - def test_wc_wildcard2_json(self): # simple json wildcard with default options - self.process( - PromptPair("the choices are: __json/wildcard2__", ""), - PromptPair("the choices are: choice3-choice1", ""), - ppp="nocup", - ) - - def test_wc_wildcard2_yaml(self): # simple yaml wildcard with default options - self.process( - PromptPair("the choices are: __yaml/wildcard2__", ""), - PromptPair("the choices are: choice3-choice1", ""), - ppp="nocup", - ) - - def test_wc_test2_yaml(self): # simple yaml wildcard - self.process( - PromptPair("the choice is: __testwc/test2__", ""), - PromptPair("the choice is: 2", ""), - ppp="nocup", - ) - - def test_wc_test3_yaml(self): # simple yaml wildcard - self.process( - PromptPair("the choice is: __testwc/test3__", ""), - PromptPair("the choice is: one choice", ""), - ppp="nocup", - ) - - def test_wc_wildcard_filter_index(self): # wildcard with positional index filter - self.process( - PromptPair("the choice is: __yaml/wildcard2'2'__", ""), - PromptPair("the choice is: choice3-choice3", ""), - ppp="nocup", - ) - - def test_wc_wildcard_filter_label(self): # wildcard with label filter - self.process( - PromptPair("the choice is: __yaml/wildcard2'label1'__", ""), - PromptPair("the choice is: choice3-choice1", ""), - ppp="nocup", - ) - - def test_wc_wildcard_filter_label2(self): # wildcard with label filter in multiple choices - self.process( - PromptPair("the choice is: __yaml/wildcard2'label2'__", ""), - PromptPair("the choice is: choice1-choice1", ""), - ppp="nocup", - ) - - def test_wc_wildcard_filter_label3(self): # wildcard with multiple label filter - self.process( - PromptPair("the choice is: __yaml/wildcard2'label1,label2'__", ""), - PromptPair("the choice is: choice3-choice1", ""), - ppp="nocup", - ) - - def test_wc_wildcard_filter_indexlabel(self): # wildcard with mixed index and label filter - self.process( - PromptPair("the choice is: __yaml/wildcard2'2,label2'__", ""), - PromptPair("the choice is: choice3-choice1", ""), - ppp="nocup", - ) - - def test_wc_wildcard_filter_compound(self): # wildcard with compound filter - self.process( - PromptPair("the choice is: __yaml/wildcard2'label1+label3'__", ""), - PromptPair("the choice is: choice3-choice3", ""), - ppp="nocup", - ) - - def test_wc_wildcard_filter_compound2(self): # wildcard with inherited compound filter - self.process( - PromptPair("the choice is: __yaml/wildcard2bis'#label1+label3'__", ""), - PromptPair("the choice is: choice3bis", ""), - ppp="nocup", - ) - - def test_wc_wildcard_filter_compound3(self): # wildcard with doubly inherited compound filter - self.process( - PromptPair("the choice is: __yaml/wildcard2bisbis'#label1+label3'__", ""), - PromptPair("the choice is: choice3bisbis", ""), - ppp="nocup", - ) - - def test_wc_wildcard_default_filter(self): # wildcard with default filter - self.process( - PromptPair( - "the choice is: __yaml/wildcard2__, __yaml/wildcard2__", - "", - ), - PromptPair("the choice is: choice3-choice3, choice3-choice1- choice2 ", ""), - ppp="nocup", - ) - - def test_wc_nested_wildcard_text(self): # nested text wildcard with repeating multiple choices - self.process( - PromptPair("the choices are: __r3$$-$$text/wildcard3__", ""), - PromptPair("the choices are: choice3,choice1- choice2 ,choice3", ""), - ppp="nocup", - ) - - def test_wc_nested_wildcard_json(self): # nested json wildcard with repeating multiple choices - self.process( - PromptPair("the choices are: __r3$$-$$json/wildcard3__", ""), - PromptPair("the choices are: choice3,choice1- choice2 ,choice3", ""), - ppp="nocup", - ) - - def test_wc_nested_wildcard_yaml(self): # nested yaml wildcard with repeating multiple choices - self.process( - PromptPair("the choices are: __r3$$-$$yaml/wildcard3__", ""), - PromptPair("the choices are: choice3,choice1- choice2 ,choice3", ""), - ppp="nocup", - ) - - def test_wc_wildcard_optional(self): # empty wildcard with no error - self.process( - PromptPair("the choices are: __yaml/empty_wildcard__", ""), - PromptPair("the choices are: ", ""), - ppp="nocup", - ) - - def test_wc_wildcard4_yaml(self): # simple yaml wildcard with one option - self.process( - PromptPair("the choices are: __yaml/wildcard4__", ""), - PromptPair("the choices are: inline text", ""), - ppp="nocup", - ) - - def test_wc_wildcard6_yaml(self): # simple yaml wildcard with object formatted choices - self.process( - PromptPair("the choices are: __yaml/wildcard6__", ""), - PromptPair("the choices are: choice2", ""), - ppp="nocup", - ) - - def test_wc_choice_wildcard_mix(self): # choices with wildcard mix - self.process( - PromptPair("the choices are: {__~2$$yaml/wildcard2__|choice0}", ""), - [ - PromptPair("the choices are: choice0", ""), - PromptPair("the choices are: choice1, choice3", ""), - PromptPair("the choices are: choice1, choice3", ""), - ], - ppp="nocup", - ) - - def test_wc_unsupportedsampler(self): # unsupported sampler - self.process( - PromptPair("the choices are: __@yaml/wildcard2__", ""), - PromptPair("", ""), - ppp="nocup", - interrupted=True, - ) - - def test_wc_wildcard_globbing(self): # wildcard with globbing - self.process( - PromptPair("the choices are: __yaml/wildcard[12]__, __yaml/wildcard?__", ""), - PromptPair("the choices are: choice3-choice2, - choice2 -choice3", ""), - ppp="nocup", - ) - - def test_wc_wildcardwithvar(self): # wildcard with inline variable - self.process( - PromptPair("the choices are: __yaml/wildcard5(var=test)__, __yaml/wildcard5__", ""), - PromptPair("the choices are: inline test, inline default", ""), - ppp="nocup", - ) - - def test_wc_wildcardPS_yaml(self): # yaml wildcard with object formatted choices and options and prefix and suffix - self.process( - PromptPair("the choices are: __yaml/wildcardPS__", ""), - PromptPair("the choices are: prefix-choice2/choice3-suffix", ""), - ppp="nocup", - ) - - def test_wc_anonymouswildcard_yaml(self): # yaml anonymous wildcard - self.process( - PromptPair("the choices are: __yaml/anonwildcards__", ""), - PromptPair("the choices are: six", ""), - ppp="nocup", - ) - - def test_wc_wildcard_input(self): # simple yaml wildcard input - self.process( - PromptPair("the choices are: __yaml_input/wildcardI__", ""), - PromptPair("the choices are: choice2", ""), - ppp="nocup", - ) - - def test_wc_circular(self): # wildcard circular reference - self.process( - PromptPair("the choices are: __yaml/circular1__", ""), - PromptPair("", ""), - ppp="nocup", - interrupted=True, - ) - - def test_wc_including(self): # wildcard including another wildcard - self.process( - PromptPair("the choices are: __yaml/including__", ""), - PromptPair("the choices are: choice4", ""), - ppp="nocup", - ) - - def test_wc_circular_including(self): # wildcard including another wildcard in a circular reference - self.process( - PromptPair("the choices are: __yaml/including1__", ""), - PromptPair("", ""), - ppp="nocup", - interrupted=True, - ) - - def test_wc_dynamicwildcard(self): # wildcard built from variables - self.process( - PromptPair( - "the choices are: ${x={1|2|3}}${w=yaml/wildcard${x}}__yaml/wildcard${x}__ __${w}__ ____", - "", - ), - PromptPair("the choices are: choice1-choice3-choice1 choice3- choice2 - choice2 choice3", ""), - ppp="nocup", - ) - - # Hosts tests - - def test_host_attention_parentheses(self): - self.process( - PromptPair( - "[test1] (test2) (test3:1.5)", - "", - ), - PromptPair("(test1:0.9) (test2) (test3:1.5)", ""), - ppp=PromptPostProcessor( - self.ppp_logger, - self.interrupt, - { - **self.def_env_info, - "ppp_config": {"hosts": {"tests": {"attention": "parentheses"}}}, - }, - self.defopts, - self.grammar_content, - self.wildcards_obj, - self.extranetwork_maps_obj, - ), - ) - - def test_host_attention_disable(self): - self.process( - PromptPair( - "[test1] (test2) (test3:1.5)", - "", - ), - 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.wildcards_obj, - self.extranetwork_maps_obj, - ), - ) - - def test_host_attention_remove(self): - self.process( - PromptPair( - "[test1] (test2) (test3:1.5)", - "", - ), - PromptPair("", ""), - ppp=PromptPostProcessor( - self.ppp_logger, - self.interrupt, - { - **self.def_env_info, - "ppp_config": {"hosts": {"tests": {"attention": "remove"}}}, - }, - self.defopts, - self.grammar_content, - self.wildcards_obj, - self.extranetwork_maps_obj, - ), - ) - - def test_host_attention_error(self): - self.process( - PromptPair( - "[test1] (test2) (test3:1.5)", - "", - ), - PromptPair("", ""), - ppp=PromptPostProcessor( - self.ppp_logger, - self.interrupt, - { - **self.def_env_info, - "ppp_config": {"hosts": {"tests": {"attention": "error"}}}, - }, - self.defopts, - self.grammar_content, - self.wildcards_obj, - self.extranetwork_maps_obj, - ), - interrupted=True, - ) - - def test_host_scheduling_before(self): - self.process( - PromptPair( - "[test1:test2:0.5]", - "", - ), - 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.wildcards_obj, - self.extranetwork_maps_obj, - ), - ) - - def test_host_scheduling_after(self): - self.process( - PromptPair( - "[test1:test2:0.5]", - "", - ), - 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.wildcards_obj, - self.extranetwork_maps_obj, - ), - ) - - def test_host_scheduling_first(self): - self.process( - PromptPair( - "[test1::0.5] [:test2:0.5] [test3:test4:0.5]", - "", - ), - 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.wildcards_obj, - self.extranetwork_maps_obj, - ), - ) - - def test_host_scheduling_remove(self): - self.process( - PromptPair( - "[test1:test2:0.5]", - "", - ), - PromptPair("", ""), - ppp=PromptPostProcessor( - self.ppp_logger, - self.interrupt, - { - **self.def_env_info, - "ppp_config": {"hosts": {"tests": {"scheduling": "remove"}}}, - }, - self.defopts, - self.grammar_content, - self.wildcards_obj, - self.extranetwork_maps_obj, - ), - ) - - def test_host_scheduling_error(self): - self.process( - PromptPair( - "[test1:test2:0.5]", - "", - ), - PromptPair("", ""), - ppp=PromptPostProcessor( - self.ppp_logger, - self.interrupt, - { - **self.def_env_info, - "ppp_config": {"hosts": {"tests": {"scheduling": "error"}}}, - }, - self.defopts, - self.grammar_content, - self.wildcards_obj, - self.extranetwork_maps_obj, - ), - interrupted=True, - ) - - def test_host_alternation_first(self): - self.process( - PromptPair( - "[test1|test2|test3]", - "", - ), - 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.wildcards_obj, - self.extranetwork_maps_obj, - ), - ) - - def test_host_alternation_remove(self): - self.process( - PromptPair( - "[test1|test2|test3]", - "", - ), - PromptPair("", ""), - ppp=PromptPostProcessor( - self.ppp_logger, - self.interrupt, - { - **self.def_env_info, - "ppp_config": {"hosts": {"tests": {"alternation": "remove"}}}, - }, - self.defopts, - self.grammar_content, - self.wildcards_obj, - self.extranetwork_maps_obj, - ), - ) - - def test_host_alternation_error(self): - self.process( - PromptPair( - "[test1|test2|test3]", - "", - ), - PromptPair("", ""), - ppp=PromptPostProcessor( - self.ppp_logger, - self.interrupt, - { - **self.def_env_info, - "ppp_config": {"hosts": {"tests": {"alternation": "error"}}}, - }, - self.defopts, - self.grammar_content, - self.wildcards_obj, - self.extranetwork_maps_obj, - ), - interrupted=True, - ) - - def test_host_and_eol(self): - self.process( - PromptPair( - "test1 AND test2:2", - "", - ), - 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.wildcards_obj, - self.extranetwork_maps_obj, - ), - ) - - def test_host_and_comma(self): - self.process( - PromptPair( - "test1 AND test2:2", - "", - ), - 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.wildcards_obj, - self.extranetwork_maps_obj, - ), - ) - - def test_host_and_remove(self): - self.process( - PromptPair( - "test1 AND test2:2", - "", - ), - 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.wildcards_obj, - self.extranetwork_maps_obj, - ), - ) - - def test_host_and_error(self): - self.process( - PromptPair( - "test1 AND test2:2", - "", - ), - PromptPair("", ""), - ppp=PromptPostProcessor( - self.ppp_logger, - self.interrupt, - { - **self.def_env_info, - "ppp_config": {"hosts": {"tests": {"and": "error"}}}, - }, - self.defopts, - self.grammar_content, - self.wildcards_obj, - self.extranetwork_maps_obj, - ), - interrupted=True, - ) - - def test_host_break_eol(self): - self.process( - PromptPair( - "test1 BREAK test2", - "", - ), - 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.wildcards_obj, - self.extranetwork_maps_obj, - ), - ) - - def test_host_break_comma(self): - self.process( - PromptPair( - "test1 BREAK test2", - "", - ), - 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.wildcards_obj, - self.extranetwork_maps_obj, - ), - ) - - def test_host_break_remove(self): - self.process( - PromptPair( - "test1 BREAK test2", - "", - ), - 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.wildcards_obj, - self.extranetwork_maps_obj, - ), - ) - - def test_host_break_error(self): - self.process( - PromptPair( - "test1 BREAK test2", - "", - ), - PromptPair("", ""), - ppp=PromptPostProcessor( - self.ppp_logger, - self.interrupt, - { - **self.def_env_info, - "ppp_config": {"hosts": {"tests": {"break": "error"}}}, - }, - self.defopts, - self.grammar_content, - self.wildcards_obj, - self.extranetwork_maps_obj, - ), - interrupted=True, - ) - - # Model variants tests - - def test_variants(self): - self.process( - PromptPair( - "test1test2test3test4", - "", - ), - PromptPair("test1test2", ""), - ppp=PromptPostProcessor( - self.ppp_logger, - self.interrupt, - { - **self.def_env_info, - "model_filename": "./webui/models/Stable-diffusion/testmodel.safetensors", - "ppp_config": { - "models": { - "sd1": { - "detect": {"tests": {"class": ["SD15", "SD15_instructpix2pix"]}}, - "variants": { - "test3": {"find_in_filename": "testmodel"}, - "sdxl": {"find_in_filename": "testmodel"}, - }, - }, - "sdxl": { - "detect": { - "tests": { - "class": [ - "SDXL", - "SDXLRefiner", - "SDXL_instructpix2pix", - "Segmind_Vega", - "KOALA_700M", - "KOALA_1B", - ] - } - }, - "variants": { - "test1": {"find_in_filename": "testmodel"}, - "test2": {"find_in_filename": "testmodel"}, - }, - }, - "something": { - "detect": {"tests": {"class": ["something"]}}, - "variants": { - "test4": {"find_in_filename": "testmodel"}, - }, - }, - } - }, - }, - { - **self.defopts, - "on_warning": PromptPostProcessor.ONWARNING_CHOICES.warn.value, - }, - self.grammar_content, - self.wildcards_obj, - self.extranetwork_maps_obj, - ), - ) - - # ComfyUI tests - - def test_comfyui_attention(self): # attention conversion - self.process( - PromptPair("(test1) (test2:1.5) [test3] [(test4)]", ""), - PromptPair("(test1) (test2:1.5) (test3:0.9) (test4:0.99)", ""), - ppp="comfyui", - ) - - # Performance tests - - def test_parser_performance_simple_simpleparser( - self, - ): # performance test with a large prompt without new constructs - large_prompt = ", ".join( - ["(this:1.2) is a [test] using a [simple|low complexity] prompt with "] * 15 - ) - self.process( - PromptPair(large_prompt, ""), - ppp="nocup", - ) - - def test_parser_performance_simple_fullparser( - self, - ): # performance test with a large prompt without new constructs but using full parser - # we trick it to use the full parser by including some characters - large_prompt = "{__${x:}}" + ", ".join( - ["(this:1.2) is a [test] using a [simple|low complexity] prompt with "] * 15 - ) - self.process( - PromptPair(large_prompt, ""), - ppp="nocup", - ) - - def test_parser_performance_complex_fullparser( - self, - ): # performance test with a large prompt with new constructs (full parser) - large_prompt = ", ".join(["__yaml/wildcard1__, (__yaml/wildcard2__), __yaml/wildcard3__, {one|two|three}"] * 15) - self.process( - PromptPair(large_prompt, ""), - ppp="nocup", - ) - - # the following tests are performance tests with only one kind of the old constructs - # same number of constructs and approximately the same full length - - def test_parser_performance_simple_attention(self): # performance test with only attention - large_prompt = ", ".join(["(one:1.2) two (three) four [five] six"] * 20) - self.process( - PromptPair(large_prompt, ""), - ppp="nocup", - ) - - def test_parser_performance_simple_schedules(self): # performance test with only schedules - large_prompt = ", ".join(["[one:1:0.5] two [three:0.8] four [five:5:0.2] six"] * 20) - self.process( - PromptPair(large_prompt, ""), - ppp="nocup", - ) - - def test_parser_performance_simple_alternation(self): # performance test with only alternation - large_prompt = ", ".join(["[one|1] two [three|3] four [five|5] six"] * 20) - self.process( - PromptPair(large_prompt, ""), - ppp="nocup", - ) - - def test_parser_performance_simple_extranetwork(self): # performance test with only extra networks - large_prompt = ", ".join([" two four six"] * 20) - self.process( - PromptPair(large_prompt, ""), - ppp="nocup", - ) - - # Variable-vs-variable comparison tests - - def test_cmd_if_var_vs_var_eq(self): # var eq var: both set to same value, if-branch taken - self.process( - PromptPair( - "hellohelloYESNO", - "", - ), - PromptPair("YES", ""), - ) - - def test_cmd_if_var_vs_var_ne(self): # var ne var: different values, ne condition true - self.process( - PromptPair( - "appleorangeYESNO", - "", - ), - PromptPair("YES", ""), - ) - - def test_cmd_if_var_vs_var_contains(self): # var contains var: var1 contains var2's value - self.process( - PromptPair( - "hello worldhelloYESNO", - "", - ), - PromptPair("YES", ""), - ) - - def test_cmd_if_var_vs_var_not_contains(self): # var not contains var: var1 does not contain var2's value - self.process( - PromptPair( - "hello worldgoodbyeYESNO", - "", - ), - PromptPair("YES", ""), - ) - - # NaN/undefined variable integer comparison tests - - def test_cmd_if_undefined_var_int_compare_warn(self): # undefined var integer compare with on_warning=warn - self.process( - PromptPair( - "YESNO", - "", - ), - PromptPair("NO", ""), - ppp=PromptPostProcessor( - self.ppp_logger, - self.interrupt, - self.def_env_info, - {**self.defopts, "on_warning": PromptPostProcessor.ONWARNING_CHOICES.warn.value}, - self.grammar_content, - self.wildcards_obj, - self.extranetwork_maps_obj, - ), - ) - - def test_cmd_if_undefined_var_int_compare_stop(self): # undefined var integer compare with on_warning=stop - self.process( - PromptPair( - "YESNO", - "", - ), - PromptPair("", ""), - interrupted=True, - ) - - def test_cmd_if_nonnumeric_var_int_compare_warn(self): # non-numeric var integer compare with on_warning=warn - self.process( - PromptPair( - "abcYESNO", - "", - ), - PromptPair("NO", ""), - ppp=PromptPostProcessor( - self.ppp_logger, - self.interrupt, - self.def_env_info, - {**self.defopts, "on_warning": PromptPostProcessor.ONWARNING_CHOICES.warn.value}, - self.grammar_content, - self.wildcards_obj, - self.extranetwork_maps_obj, - ), - ) - - def test_cmd_if_nonnumeric_var_int_compare_stop(self): # non-numeric var integer compare with on_warning=stop - self.process( - PromptPair( - "abcYESNO", - "", - ), - PromptPair("", ""), - interrupted=True, - ) - - def test_cmd_if_empty_var_int_compare(self): # empty string var integer compare with on_warning=warn - self.process( - PromptPair( - "YESNO", - "", - ), - PromptPair("NO", ""), - ppp=PromptPostProcessor( - self.ppp_logger, - self.interrupt, - self.def_env_info, - {**self.defopts, "on_warning": PromptPostProcessor.ONWARNING_CHOICES.warn.value}, - self.grammar_content, - self.wildcards_obj, - self.extranetwork_maps_obj, - ), - ) - - -if __name__ == "__main__": - unittest.main() diff --git a/tests/tests_choices.py b/tests/tests_choices.py new file mode 100644 index 0000000..2d8f6ce --- /dev/null +++ b/tests/tests_choices.py @@ -0,0 +1,101 @@ +import unittest + +from ppp import PromptPostProcessor # pylint: disable=import-error +from .base_tests import PromptPair, TestPromptPostProcessorBase + + +class TestChoices(TestPromptPostProcessorBase): + + def setUp(self): # pylint: disable=arguments-differ + super().setUp(enable_file_logging=False) + + # Choices tests + + def test_ch_choices(self): # simple choices with weights + self.process( + PromptPair("the choices are: {3::choice1|2::choice2|choice3}", ""), + PromptPair("the choices are: choice2", ""), + ppp="nocup", + ) + + def test_ch_unsupportedsampler(self): # unsupported sampler + self.process( + PromptPair("the choices are: {@choice1|choice2|choice3}", ""), + PromptPair("", ""), + ppp="nocup", + interrupted=True, + ) + + def test_ch_choices_withcomments(self): # choices with comments and multiline + self.process( + PromptPair( + "the choices are: {\n3::choice1 # this is option 1\n|2::choice2\n# this was option 2\n|choice3 # this is option 3\n}", + "", + ), + PromptPair("the choices are: choice2", ""), + ppp="nocup", + ) + + def test_ch_choices_multiple(self): # choices with multiple selection + self.process( + PromptPair("the choices are: {~2$$, $$3::choice1|2:: choice2 |choice3}", ""), + PromptPair("the choices are: choice2 , choice3", ""), + ppp="nocup", + ) + + def test_ch_choices_if_multiple(self): # choices with if and multiple selection + self.process( + PromptPair("the choices are: {2$$, $$3::choice1|2 if _is_sd1::choice2|choice3}", ""), + PromptPair("the choices are: choice1, choice3", ""), + ppp="nocup", + ) + + def test_ch_choices_set_if_multiple(self): # choices with if user variable and multiple selection + self.process( + PromptPair("${var=test}the choices are: {2$$, $$3::choice1|2 if not var eq 'test'::choice2|choice3}", ""), + PromptPair("the choices are: choice1, choice3", ""), + ppp="nocup", + ) + + def test_ch_choices_set_if_nested(self): # nested choices with if user variable and multiple selection + self.process( + PromptPair( + "${var=test}the choices are: {2$$, $$3::choice1${var2=test2} {if var2 eq 'test2'::choice11|choice12}|2 if not var eq 'test'::choice2|choice3}", + "", + ), + PromptPair("the choices are: choice1 choice11, choice3", ""), + ppp="nocup", + ) + + def test_ch_choicesinsidelora(self): # simple choices inside a lora + self.process( + PromptPair("", ""), + PromptPair("", ""), + ppp="nocup", + ) + + def test_ch_removelorawithchoices(self): + self.process( + PromptPair("", ""), + PromptPair("", ""), + ppp=PromptPostProcessor( + self.ppp_logger, + self.interrupt, + self.def_env_info, + {**self.defopts, "remove_extranetwork_tags": True}, + self.grammar_content, + self.wildcards_obj, + self.extranetwork_maps_obj, + ), + ) + + def test_ch_cmd_includewildcard(self): + self.process( + PromptPair("{ch_one|ch_two|%0.5::include yaml/wildcard1}", ""), + PromptPair("ch_two", ""), + ppp="nocup", + ) + + +if __name__ == "__main__": + unittest.main() diff --git a/tests/tests_cleanup.py b/tests/tests_cleanup.py new file mode 100644 index 0000000..c6472ae --- /dev/null +++ b/tests/tests_cleanup.py @@ -0,0 +1,132 @@ +import unittest + +from ppp import PromptPostProcessor # pylint: disable=import-error + +from .base_tests import PromptPair, TestPromptPostProcessorBase + + +class TestCleanup(TestPromptPostProcessorBase): + + def setUp(self): # pylint: disable=arguments-differ + super().setUp(enable_file_logging=False) + + # Cleanup tests + + def test_cl_simple(self): # simple cleanup + self.process( + PromptPair(" this is a ((test ), , , (), , [] ( , test ,:2.0):1.5), (red:1.5) ", " normal quality "), + PromptPair("this is a ((test), (test,:2):1.5), (red:1.5)", "normal quality"), + ) + + def test_cl_complex(self): # complex cleanup + self.process( + PromptPair( + " this is BREAKABLE a ((test)), ,AND AND(() [] ANDERSON (test:2.0):1.5) :o BREAK \n BREAK (red:1.5) ", + " [:hands, feet, :0.15]normal quality ", + ), + PromptPair( + "this is BREAKABLE a (test:1.21) AND( ANDERSON (test:2):1.5) :o BREAK (red:1.5)", + "[:hands, feet, :0.15]normal quality", + ), + ) + + def test_cl_removenetworktags(self): # remove network tags + self.process( + PromptPair("this is a test__yaml/wildcard7__", ""), + PromptPair("this is a test", ""), + ppp=PromptPostProcessor( + self.ppp_logger, + self.interrupt, + self.def_env_info, + {**self.defopts, "remove_extranetwork_tags": True}, + self.grammar_content, + self.wildcards_obj, + self.extranetwork_maps_obj, + ), + ) + + def test_cl_dontremoveseparatorsoneol(self): # don't remove separators on eol + self.process( + PromptPair("this is a test,\nsecond line", ""), + 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, + }, + self.grammar_content, + self.wildcards_obj, + self.extranetwork_maps_obj, + ), + ) + + def test_cl_separatorswitheol(self): # don't remove eols with the separators + self.process( + PromptPair( + """{ (d:0.9) ,, (l:1.1) | (l:1.1) (d:0.9),,, } + (l:1.1) + (d:0.9)""", + "", + ), + PromptPair( + """ (l:1.1) (d:0.9), + (l:1.1) + (d:0.9)""", + "", + ), + 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, + }, + self.grammar_content, + self.wildcards_obj, + self.extranetwork_maps_obj, + ), + ) + + def test_cl_mergeattention(self): # merge attention + self.process( + PromptPair( + "this is (a test:0.9) of (attention (merging:1.2)) where ((this)) ((is joined:1.2)) and ([this too]:1.3)", + "", + ), + PromptPair( + "this is [a test] of (attention (merging:1.2)) where (this:1.21) (is joined:1.32) and (this too:1.17)", + "", + ), + ) + + def test_cl_not_mergeattention(self): # not merge attention + self.process( + PromptPair( + "this is (a test:0.9) of not (attention (merging:1.2)) where ((this)) ((is not joined:1.2)) and neither is ([this]:1.3)", + "", + ), + PromptPair( + "this is (a test:0.9) of not (attention (merging:1.2)) where ((this)) ((is not joined:1.2)) and neither is ([this]:1.3)", + "", + ), + ppp="nocup", + ) + + +if __name__ == "__main__": + unittest.main() diff --git a/tests/tests_commands.py b/tests/tests_commands.py new file mode 100644 index 0000000..edd0eec --- /dev/null +++ b/tests/tests_commands.py @@ -0,0 +1,348 @@ +import unittest + +from ppp import PromptPostProcessor # pylint: disable=import-error +from .base_tests import PromptPair, TestPromptPostProcessorBase + + +class TestCommands(TestPromptPostProcessorBase): + + def setUp(self): # pylint: disable=arguments-differ + super().setUp(enable_file_logging=False) + + # Command tests + + def test_cmd_stn_complex_features(self): # complex stn command with AND, BREAK and other features + self.process( + PromptPair( + "[neg5] this \\(is\\): a (([complex|simpleneg6|regular] testneg1)(test:2.0):1.5) \nBREAK, BREAK with [abcneg4:defneg2(neg3:1.6):5]:0.5 AND loratrigger AND AND hypernettrigger :0.3", + "normal quality, ", + ), + PromptPair( + "this \\(is\\): a (([complex|simple|regular] test)(test:2):1.5)\nBREAK with [abc:def:5]:0.5 AND loratrigger AND hypernettrigger :0.3", + "[neg5], ([|neg6|]:1.65), (neg1:1.65), [neg4::5], normal quality, [neg2(neg3:1.6):5]", + ), + ) + + def test_cmd_if_complex_features(self): # complex if command + self.process( + PromptPair( + "this \\(is\\): a (([complex|simple|regular] test)(test:2.0):1.5) \nBREAK, BREAK with [abcneg4:def:5]:0.5 AND loratrigger hypernettrigger nothing:0.3", + "normal quality", + ), + PromptPair( + "this \\(is\\): a (([complex|simple|regular] test)(test:2):1.5)\nBREAK :0.5 AND hypernettrigger :0.3", + "normal quality", + ), + ) + + def test_cmd_if_nested(self): # nested if command + self.process( + PromptPair( + "this is SD1PONYSD2NOPONYNOPONY", + "", + ), + 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.wildcards_obj, + self.extranetwork_maps_obj, + ), + ) + + def test_cmd_set_if(self): # set and if commands + self.process( + PromptPair("valuethis test is OKnot OK", ""), + PromptPair("this test is OK", ""), + ) + + def test_cmd_set_empty(self): # set to empty + self.process( + PromptPair("${v2=}this test is not OKOK", ""), + PromptPair("this test is OK", ""), + ) + + def test_cmd_set_eval_if(self): # set and if commands + self.process( + PromptPair("valuethis test is OKnot OK", ""), + PromptPair("this test is OK", ""), + ) + + def test_cmd_set_if_echo_nested(self): # nested set, if and echo commands + self.process( + PromptPair( + "1OKnot OK NOK OK", + "", + ), + PromptPair("OK OK OK", ""), + ) + + def test_cmd_set_if_complex_conditions_1(self): # complex conditions (or) + self.process( + PromptPair( + "truefalsethis test is OKnot OK", + "", + ), + PromptPair("this test is OK", ""), + ) + + def test_cmd_set_if_complex_conditions_2(self): # complex conditions (and) + self.process( + PromptPair( + "truetruethis test is OKnot OK", + "", + ), + PromptPair("this test is OK", ""), + ) + + def test_cmd_set_if_complex_conditions_3(self): # complex conditions (not) + self.process( + PromptPair("falsethis test is OKnot OK", ""), + PromptPair("this test is OK", ""), + ) + + def test_cmd_set_if_complex_conditions_4(self): # complex conditions (not, precedence) + self.process( + PromptPair( + "truefalsethis test is OKnot OK", + "", + ), + PromptPair("this test is OK", ""), + ) + + def test_cmd_set_if_complex_conditions_5(self): # complex conditions (not, precedence, comparison) + self.process( + PromptPair( + "1falsethis test is OKnot OK", + "", + ), + PromptPair("this test is OK", ""), + ) + + def test_cmd_set_if_complex_conditions_6(self): # complex conditions + self.process( + PromptPair( + "123this test is OKnot OK", + "", + ), + PromptPair("this test is OK", ""), + ) + + def test_cmd_set_if_complex_conditions_7(self): # complex conditions + self.process( + PromptPair( + "123this test is OKnot OK", + "", + ), + PromptPair("this test is OK", ""), + ) + + def test_cmd_set_if2(self): # set and more complex if commands + self.process( + PromptPair( + "First: value1this test is OKOK2not OK\nSecond: value3this test is OKnot OK", + "", + ), + PromptPair("First: this test is OK\nSecond: this test is OK", ""), + ) + + def test_cmd_set_add_if(self): # set, add and if commands + self.process( + PromptPair( + "value2this test is OKnot OK", + "", + ), + PromptPair("this test is OK", ""), + ) + + def test_cmd_set_add_DP_if(self): # set, add (DP format) and if commands + self.process( + PromptPair( + "${v=value}${v+=2}this test is OKnot OK", + "", + ), + PromptPair("this test is OK", ""), + ) + + def test_cmd_set_immediateeval(self): # set (DP format) with mixed evaluation + self.process( + PromptPair( + "${var=!__yaml/wildcard1__}the choices are: ${var}, ${var}, ${var2:default}, ${var3=__yaml/wildcard1__}${var3}, ${var3}", + "", + ), + PromptPair("the choices are: choice2, choice2, default, choice3, choice1", ""), + ppp="nocup", + ) + + def test_cmd_set_mixeval(self): # set and add (DP format) with mixed evaluation + self.process( + PromptPair( + "${var=__yaml/wildcard1__}the choices are: ${var}, ${var}, ${var+=, __yaml/wildcard2__}${var}, ${var}, ${var+=!, __yaml/wildcard3__}${var}, ${var}", + "", + ), + PromptPair( + "the choices are: choice2, choice3, choice1, choice1- choice2 -choice3, choice2, choice2 -choice1-choice3, choice2, choice3-choice1- choice2 , choice1, choice2 , choice2, choice3-choice1- choice2 , choice1, choice2 ", + "", + ), + ppp="nocup", + ) + + def test_cmd_set_ifundefined_if(self): # set, ifundefined and if commands + self.process( + PromptPair( + "valuethis test is OKnot OK", + "", + ), + PromptPair("this test is OK", ""), + ) + + def test_cmd_set_ifundefined_if_2(self): # set, ifundefined and if commands + self.process( + PromptPair( + "valuevalue2this test is OKnot OK", + "", + ), + PromptPair("this test is OK", ""), + ) + + def test_cmd_set_ifundefined_DP_if(self): # set, ifundefined (DP format) and if commands + self.process( + PromptPair( + "${v?=value}this test is OKnot OK", + "", + ), + PromptPair("this test is OK", ""), + ) + + def test_cmd_set_ifundefined_DP_if_2(self): # set, ifundefined (DP format) and if commands + self.process( + PromptPair( + "${v=!value}${v?=!value2}this test is OKnot OK", + "", + ), + PromptPair("this test is OK", ""), + ) + + def test_cmd_ext(self): # ext + self.process( + PromptPair( + "trigger1trigger2trigger4trigger5", + "", + ), + PromptPair( + "trigger1,trigger2,trigger4,trigger5", + "", + ), + ) + + def test_cmd_ext_map_notrigger(self): # ext mapping, no trigger + self.process( + PromptPair( + "", + "", + ), + PromptPair("triggergeneric1, triggergeneric2, two, triggergeneric1, triggergeneric2, two", ""), + ) + + def test_cmd_ext_map1(self): # ext mapping, no lora + self.process( + PromptPair( + "inlinetrigger", + "", + ), + PromptPair("inlinetrigger, triggergeneric1, triggergeneric2, two", ""), + ) + + def test_cmd_ext_map2(self): # ext mapping, lora with weight + self.process( + PromptPair( + "inlinetrigger", + "", + ), + 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.wildcards_obj, + self.extranetwork_maps_obj, + ), + ) + + def test_cmd_ext_map3(self): # ext mapping, lora with weight adjusted + self.process( + PromptPair( + "inlinetrigger", + "", + ), + 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.wildcards_obj, + self.extranetwork_maps_obj, + ), + ) + + def test_cmd_ext_map4(self): # ext mapping, lora with parameters + self.process( + PromptPair( + "inlinetrigger", + "", + ), + 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.wildcards_obj, + self.extranetwork_maps_obj, + ), + ) + + def test_cmd_ext_map5(self): # ext mapping, lora with no parameters + self.process( + PromptPair( + "inlinetrigger", + "", + ), + 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.wildcards_obj, + self.extranetwork_maps_obj, + ), + ) + + +if __name__ == "__main__": + unittest.main() diff --git a/tests/tests_host.py b/tests/tests_host.py new file mode 100644 index 0000000..bff3fac --- /dev/null +++ b/tests/tests_host.py @@ -0,0 +1,442 @@ +import unittest + +from ppp import PromptPostProcessor # pylint: disable=import-error + +from .base_tests import PromptPair, TestPromptPostProcessorBase + + +class TestHosts(TestPromptPostProcessorBase): + + def setUp(self): # pylint: disable=arguments-differ + super().setUp(enable_file_logging=False) + + # Hosts tests + + def test_host_attention_parentheses(self): + self.process( + PromptPair( + "[test1] (test2) (test3:1.5) [(test4)]", + "", + ), + 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.wildcards_obj, + self.extranetwork_maps_obj, + ), + ) + + def test_host_attention_disable(self): + self.process( + PromptPair( + "[test1] (test2) (test3:1.5)", + "", + ), + 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.wildcards_obj, + self.extranetwork_maps_obj, + ), + ) + + def test_host_attention_remove(self): + self.process( + PromptPair( + "[test1] (test2) (test3:1.5)", + "", + ), + PromptPair("", ""), + ppp=PromptPostProcessor( + self.ppp_logger, + self.interrupt, + { + **self.def_env_info, + "ppp_config": {"hosts": {"tests": {"attention": "remove"}}}, + }, + self.defopts, + self.grammar_content, + self.wildcards_obj, + self.extranetwork_maps_obj, + ), + ) + + def test_host_attention_error(self): + self.process( + PromptPair( + "[test1] (test2) (test3:1.5)", + "", + ), + PromptPair("", ""), + ppp=PromptPostProcessor( + self.ppp_logger, + self.interrupt, + { + **self.def_env_info, + "ppp_config": {"hosts": {"tests": {"attention": "error"}}}, + }, + self.defopts, + self.grammar_content, + self.wildcards_obj, + self.extranetwork_maps_obj, + ), + interrupted=True, + ) + + def test_host_scheduling_before(self): + self.process( + PromptPair( + "[test1:test2:0.5]", + "", + ), + 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.wildcards_obj, + self.extranetwork_maps_obj, + ), + ) + + def test_host_scheduling_after(self): + self.process( + PromptPair( + "[test1:test2:0.5]", + "", + ), + 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.wildcards_obj, + self.extranetwork_maps_obj, + ), + ) + + def test_host_scheduling_first(self): + self.process( + PromptPair( + "[test1::0.5] [:test2:0.5] [test3:test4:0.5]", + "", + ), + 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.wildcards_obj, + self.extranetwork_maps_obj, + ), + ) + + def test_host_scheduling_remove(self): + self.process( + PromptPair( + "[test1:test2:0.5]", + "", + ), + PromptPair("", ""), + ppp=PromptPostProcessor( + self.ppp_logger, + self.interrupt, + { + **self.def_env_info, + "ppp_config": {"hosts": {"tests": {"scheduling": "remove"}}}, + }, + self.defopts, + self.grammar_content, + self.wildcards_obj, + self.extranetwork_maps_obj, + ), + ) + + def test_host_scheduling_error(self): + self.process( + PromptPair( + "[test1:test2:0.5]", + "", + ), + PromptPair("", ""), + ppp=PromptPostProcessor( + self.ppp_logger, + self.interrupt, + { + **self.def_env_info, + "ppp_config": {"hosts": {"tests": {"scheduling": "error"}}}, + }, + self.defopts, + self.grammar_content, + self.wildcards_obj, + self.extranetwork_maps_obj, + ), + interrupted=True, + ) + + def test_host_alternation_first(self): + self.process( + PromptPair( + "[test1|test2|test3]", + "", + ), + 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.wildcards_obj, + self.extranetwork_maps_obj, + ), + ) + + def test_host_alternation_remove(self): + self.process( + PromptPair( + "[test1|test2|test3]", + "", + ), + PromptPair("", ""), + ppp=PromptPostProcessor( + self.ppp_logger, + self.interrupt, + { + **self.def_env_info, + "ppp_config": {"hosts": {"tests": {"alternation": "remove"}}}, + }, + self.defopts, + self.grammar_content, + self.wildcards_obj, + self.extranetwork_maps_obj, + ), + ) + + def test_host_alternation_error(self): + self.process( + PromptPair( + "[test1|test2|test3]", + "", + ), + PromptPair("", ""), + ppp=PromptPostProcessor( + self.ppp_logger, + self.interrupt, + { + **self.def_env_info, + "ppp_config": {"hosts": {"tests": {"alternation": "error"}}}, + }, + self.defopts, + self.grammar_content, + self.wildcards_obj, + self.extranetwork_maps_obj, + ), + interrupted=True, + ) + + def test_host_and_eol(self): + self.process( + PromptPair( + "test1 AND test2:2", + "", + ), + 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.wildcards_obj, + self.extranetwork_maps_obj, + ), + ) + + def test_host_and_comma(self): + self.process( + PromptPair( + "test1 AND test2:2", + "", + ), + 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.wildcards_obj, + self.extranetwork_maps_obj, + ), + ) + + def test_host_and_remove(self): + self.process( + PromptPair( + "test1 AND test2:2", + "", + ), + 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.wildcards_obj, + self.extranetwork_maps_obj, + ), + ) + + def test_host_and_error(self): + self.process( + PromptPair( + "test1 AND test2:2", + "", + ), + PromptPair("", ""), + ppp=PromptPostProcessor( + self.ppp_logger, + self.interrupt, + { + **self.def_env_info, + "ppp_config": {"hosts": {"tests": {"and": "error"}}}, + }, + self.defopts, + self.grammar_content, + self.wildcards_obj, + self.extranetwork_maps_obj, + ), + interrupted=True, + ) + + def test_host_break_eol(self): + self.process( + PromptPair( + "test1 BREAK test2", + "", + ), + 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.wildcards_obj, + self.extranetwork_maps_obj, + ), + ) + + def test_host_break_comma(self): + self.process( + PromptPair( + "test1 BREAK test2", + "", + ), + 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.wildcards_obj, + self.extranetwork_maps_obj, + ), + ) + + def test_host_break_remove(self): + self.process( + PromptPair( + "test1 BREAK test2", + "", + ), + 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.wildcards_obj, + self.extranetwork_maps_obj, + ), + ) + + def test_host_break_error(self): + self.process( + PromptPair( + "test1 BREAK test2", + "", + ), + PromptPair("", ""), + ppp=PromptPostProcessor( + self.ppp_logger, + self.interrupt, + { + **self.def_env_info, + "ppp_config": {"hosts": {"tests": {"break": "error"}}}, + }, + self.defopts, + self.grammar_content, + self.wildcards_obj, + self.extranetwork_maps_obj, + ), + interrupted=True, + ) + + +if __name__ == "__main__": + unittest.main() diff --git a/tests/tests_performance.py b/tests/tests_performance.py new file mode 100644 index 0000000..1b73581 --- /dev/null +++ b/tests/tests_performance.py @@ -0,0 +1,78 @@ +import unittest + +from .base_tests import PromptPair, TestPromptPostProcessorBase + + +class TestPerformance(TestPromptPostProcessorBase): + + def setUp(self): # pylint: disable=arguments-differ + super().setUp(enable_file_logging=False) + + # Performance tests + + def test_parser_performance_simple_simpleparser( + self, + ): # performance test with a large prompt without new constructs + large_prompt = ", ".join( + ["(this:1.2) is a [test] using a [simple|low complexity] prompt with "] * 15 + ) + self.process( + PromptPair(large_prompt, ""), + ppp="nocup", + ) + + def test_parser_performance_simple_fullparser( + self, + ): # performance test with a large prompt without new constructs but using full parser + # we trick it to use the full parser by including some characters + large_prompt = "{__${x:}}" + ", ".join( + ["(this:1.2) is a [test] using a [simple|low complexity] prompt with "] * 15 + ) + self.process( + PromptPair(large_prompt, ""), + ppp="nocup", + ) + + def test_parser_performance_complex_fullparser( + self, + ): # performance test with a large prompt with new constructs (full parser) + large_prompt = ", ".join(["__yaml/wildcard1__, (__yaml/wildcard2__), __yaml/wildcard3__, {one|two|three}"] * 15) + self.process( + PromptPair(large_prompt, ""), + ppp="nocup", + ) + + # the following tests are performance tests with only one kind of the old constructs + # same number of constructs and approximately the same full length + + def test_parser_performance_simple_attention(self): # performance test with only attention + large_prompt = ", ".join(["(one:1.2) two (three) four [five] six"] * 20) + self.process( + PromptPair(large_prompt, ""), + ppp="nocup", + ) + + def test_parser_performance_simple_schedules(self): # performance test with only schedules + large_prompt = ", ".join(["[one:1:0.5] two [three:0.8] four [five:5:0.2] six"] * 20) + self.process( + PromptPair(large_prompt, ""), + ppp="nocup", + ) + + def test_parser_performance_simple_alternation(self): # performance test with only alternation + large_prompt = ", ".join(["[one|1] two [three|3] four [five|5] six"] * 20) + self.process( + PromptPair(large_prompt, ""), + ppp="nocup", + ) + + def test_parser_performance_simple_extranetwork(self): # performance test with only extra networks + large_prompt = ", ".join([" two four six"] * 20) + self.process( + PromptPair(large_prompt, ""), + ppp="nocup", + ) + + +if __name__ == "__main__": + unittest.main() diff --git a/tests/tests_stn.py b/tests/tests_stn.py new file mode 100644 index 0000000..4692f45 --- /dev/null +++ b/tests/tests_stn.py @@ -0,0 +1,127 @@ +import unittest + +from .base_tests import PromptPair, TestPromptPostProcessorBase + + +class TestSendToNegative(TestPromptPostProcessorBase): + + def setUp(self): # pylint: disable=arguments-differ + super().setUp(enable_file_logging=False) + + # Send To Negative tests + + def test_stn_simple(self): # negtags with different parameters and separations + self.process( + PromptPair( + "flowersred, green, blueyellow, purpleblack", + "normal quality, worse quality", + ), + PromptPair("flowers", "red, green, yellow, normal quality, purple, worse quality, black, blue"), + ) + + def test_stn_complex(self): # complex negtags + self.process( + PromptPair( + "red ((pink)), flowers purple, mauveblue, yellow green", + "normal quality, , bad quality, worse quality", + ), + PromptPair( + "flowers", + "red, (pink:1.21), normal quality, mauve, yellow, bad quality, green, worse quality, purple, blue", + ), + ) + + def test_stn_complex_nocleanup(self): # complex negtags with no cleanup + self.process( + PromptPair( + "red ((pink)), flowers purple, mauveblue, yellow green", + "normal quality, , bad quality, worse quality", + ), + PromptPair( + " (()), flowers , , ", + "red, ((pink)), normal quality, mauve, yellow, bad quality, green, worse quality, purple, blue", + ), + ppp="nocup", + ) + + def test_stn_inside_attention(self): # negtag inside attention + self.process( + PromptPair( + "[neg1] this is a ((testneg2) (test:2.0): 1.5 ) (red[square]:1.5)", + "normal quality", + ), + PromptPair( + "this is a ((test) (test:2):1.5) (red:1.5)", "[neg1], ([square]:1.5), normal quality, (neg2:1.65)" + ), + ) + + def test_stn_inside_alternation(self): # negtag inside alternation + self.process( + PromptPair( + "this is a (([complexneg1|simpleneg2|regularneg3] test)(test:2.0):1.5)", + "normal quality", + ), + PromptPair( + "this is a (([complex|simple|regular] test)(test:2):1.5)", + "([neg1||]:1.65), ([|neg2|]:1.65), ([||neg3]:1.65), normal quality", + ), + ) + + def test_stn_inside_alternation_recursive(self): # negtag inside alternation (recursive alternation) + self.process( + PromptPair( + "this is a (([complexneg1[one|twoneg12||three|four(neg14)]|simpleneg2|regularneg3] test)(test:2.0):1.5)", + "normal quality", + ), + PromptPair( + "this is a (([complex[one|two||three|four]|simple|regular] test)(test:2):1.5)", + "([neg1||]:1.65), ([[|neg12|||]||]:1.65), ([[||||(neg14)]||]:1.65), ([|neg2|]:1.65), ([||neg3]:1.65), normal quality", + ), + ) + + def test_stn_inside_scheduling(self): # negtag inside scheduling + self.process( + PromptPair("this is [abcneg1:defneg2: 5 ]", "normal quality"), + [PromptPair("this is [abc:def:5]", "[neg1::5], normal quality, [neg2:5]")], + ) + + def test_stn_complex_features(self): # complex negtags with AND, BREAK and other features + self.process( + PromptPair( + "[neg5] this \\(is\\): a (([complex|simpleneg6|regular] testneg1)(test:2.0):1.5) \nBREAK, BREAK with [abcneg4:defneg2(neg3:1.6):5]:0.5 AND loratrigger AND AND hypernettrigger :0.3", + "normal quality, ", + ), + PromptPair( + "this \\(is\\): a (([complex|simple|regular] test)(test:2):1.5)\nBREAK with [abc:def:5]:0.5 AND loratrigger AND hypernettrigger :0.3", + "[neg5], ([|neg6|]:1.65), (neg1:1.65), [neg4::5], normal quality, [neg2(neg3:1.6):5]", + ), + ) + + def test_stn_complex_features_newformat(self): # complex negtags with AND, BREAK and other features (new format) + self.process( + PromptPair( + "[neg5] this \\(is\\): a (([complex|simpleneg6|regular] testneg1)(test:2.0):1.5) \nBREAK, BREAK with [abcneg4:defneg2(neg3:1.6):5]:0.5 AND loratrigger AND AND hypernettrigger :0.3", + "normal quality, ", + ), + PromptPair( + "this \\(is\\): a (([complex|simple|regular] test)(test:2):1.5)\nBREAK with [abc:def:5]:0.5 AND loratrigger AND hypernettrigger :0.3", + "[neg5], ([|neg6|]:1.65), (neg1:1.65), [neg4::5], normal quality, [neg2(neg3:1.6):5]", + ), + ) + + def test_stn_inside_alternation_recursive_2(self): # negtag inside alternation (recursive alternation) + self.process( + PromptPair( + "[pos1neg1[pos11|pos12neg12||pos14|pos15neg15]|pos2neg2|pos3neg3]", + "", + ), + PromptPair( + "[pos1[pos11|pos12||pos14|pos15]|pos2|pos3]", + "[neg1||], [[|neg12|||]||], [[||||neg15]||], [|neg2|], [||neg3]", + # "[neg1[|neg12|||neg15]|neg2|neg3]", # expected output if the constructs were unified + ), + ) + + +if __name__ == "__main__": + unittest.main() diff --git a/tests/tests_variables.py b/tests/tests_variables.py new file mode 100644 index 0000000..1571d08 --- /dev/null +++ b/tests/tests_variables.py @@ -0,0 +1,160 @@ +import unittest + +from ppp import PromptPostProcessor # pylint: disable=import-error +from .base_tests import PromptPair, TestPromptPostProcessorBase + + +class TestVariables(TestPromptPostProcessorBase): + + def setUp(self): # pylint: disable=arguments-differ + super().setUp(enable_file_logging=False) + + # Variable nesting tests + + def test_var_nested_1(self): # variable default nested in variable set + self.process( + PromptPair( + "${v1=test ${v2:OK}}${v1}", + "", + ), + PromptPair("test OK", ""), + variables={"v1": "test OK", "v2": "OK"}, + ) + + def test_var_nested_2(self): # variable set nested in variable default + self.process( + PromptPair( + "${v1:test ${v2=OK}${v2}}", + "", + ), + PromptPair("test OK", ""), + variables={"v1": "test OK", "v2": "OK"}, + ) + + def test_var_nested_3(self): # variable default nested in variable default + self.process( + PromptPair( + "${v1:test ${v2:OK}}", + "", + ), + PromptPair("test OK", ""), + variables={"v1": "test OK", "v2": "OK"}, + ) + + # Variable-vs-variable comparison tests + + def test_cmd_if_var_vs_var_eq(self): # var eq var: both set to same value, if-branch taken + self.process( + PromptPair( + "hellohelloYESNO", + "", + ), + PromptPair("YES", ""), + ) + + def test_cmd_if_var_vs_var_ne(self): # var ne var: different values, ne condition true + self.process( + PromptPair( + "appleorangeYESNO", + "", + ), + PromptPair("YES", ""), + ) + + def test_cmd_if_var_vs_var_contains(self): # var contains var: var1 contains var2's value + self.process( + PromptPair( + "hello worldhelloYESNO", + "", + ), + PromptPair("YES", ""), + ) + + def test_cmd_if_var_vs_var_not_contains(self): # var not contains var: var1 does not contain var2's value + self.process( + PromptPair( + "hello worldgoodbyeYESNO", + "", + ), + PromptPair("YES", ""), + ) + + # NaN/undefined variable integer comparison tests + + def test_cmd_if_undefined_var_int_compare_warn(self): # undefined var integer compare with on_warning=warn + self.process( + PromptPair( + "YESNO", + "", + ), + PromptPair("NO", ""), + ppp=PromptPostProcessor( + self.ppp_logger, + self.interrupt, + self.def_env_info, + {**self.defopts, "on_warning": PromptPostProcessor.ONWARNING_CHOICES.warn.value}, + self.grammar_content, + self.wildcards_obj, + self.extranetwork_maps_obj, + ), + ) + + def test_cmd_if_undefined_var_int_compare_stop(self): # undefined var integer compare with on_warning=stop + self.process( + PromptPair( + "YESNO", + "", + ), + PromptPair("", ""), + interrupted=True, + ) + + def test_cmd_if_nonnumeric_var_int_compare_warn(self): # non-numeric var integer compare with on_warning=warn + self.process( + PromptPair( + "abcYESNO", + "", + ), + PromptPair("NO", ""), + ppp=PromptPostProcessor( + self.ppp_logger, + self.interrupt, + self.def_env_info, + {**self.defopts, "on_warning": PromptPostProcessor.ONWARNING_CHOICES.warn.value}, + self.grammar_content, + self.wildcards_obj, + self.extranetwork_maps_obj, + ), + ) + + def test_cmd_if_nonnumeric_var_int_compare_stop(self): # non-numeric var integer compare with on_warning=stop + self.process( + PromptPair( + "abcYESNO", + "", + ), + PromptPair("", ""), + interrupted=True, + ) + + def test_cmd_if_empty_var_int_compare(self): # empty string var integer compare with on_warning=warn + self.process( + PromptPair( + "YESNO", + "", + ), + PromptPair("NO", ""), + ppp=PromptPostProcessor( + self.ppp_logger, + self.interrupt, + self.def_env_info, + {**self.defopts, "on_warning": PromptPostProcessor.ONWARNING_CHOICES.warn.value}, + self.grammar_content, + self.wildcards_obj, + self.extranetwork_maps_obj, + ), + ) + + +if __name__ == "__main__": + unittest.main() diff --git a/tests/tests_variants.py b/tests/tests_variants.py new file mode 100644 index 0000000..8ec1bc0 --- /dev/null +++ b/tests/tests_variants.py @@ -0,0 +1,75 @@ +import unittest + +from ppp import PromptPostProcessor # pylint: disable=import-error +from .base_tests import PromptPair, TestPromptPostProcessorBase + + +class TestModelVariants(TestPromptPostProcessorBase): + + def setUp(self): # pylint: disable=arguments-differ + super().setUp(enable_file_logging=False) + + # Model variants tests + + def test_variants(self): + self.process( + PromptPair( + "test1test2test3test4", + "", + ), + PromptPair("test1test2", ""), + ppp=PromptPostProcessor( + self.ppp_logger, + self.interrupt, + { + **self.def_env_info, + "model_filename": "./webui/models/Stable-diffusion/testmodel.safetensors", + "ppp_config": { + "models": { + "sd1": { + "detect": {"tests": {"class": ["SD15", "SD15_instructpix2pix"]}}, + "variants": { + "test3": {"find_in_filename": "testmodel"}, + "sdxl": {"find_in_filename": "testmodel"}, + }, + }, + "sdxl": { + "detect": { + "tests": { + "class": [ + "SDXL", + "SDXLRefiner", + "SDXL_instructpix2pix", + "Segmind_Vega", + "KOALA_700M", + "KOALA_1B", + ] + } + }, + "variants": { + "test1": {"find_in_filename": "testmodel"}, + "test2": {"find_in_filename": "testmodel"}, + }, + }, + "something": { + "detect": {"tests": {"class": ["something"]}}, + "variants": { + "test4": {"find_in_filename": "testmodel"}, + }, + }, + } + }, + }, + { + **self.defopts, + "on_warning": PromptPostProcessor.ONWARNING_CHOICES.warn.value, + }, + self.grammar_content, + self.wildcards_obj, + self.extranetwork_maps_obj, + ), + ) + + +if __name__ == "__main__": + unittest.main() diff --git a/tests/tests_wildcards.py b/tests/tests_wildcards.py new file mode 100644 index 0000000..343bfa8 --- /dev/null +++ b/tests/tests_wildcards.py @@ -0,0 +1,400 @@ +import unittest + +from ppp import PromptPostProcessor +from .base_tests import PromptPair, TestPromptPostProcessorBase + + +class TestWildcards(TestPromptPostProcessorBase): + + def setUp(self): # pylint: disable=arguments-differ + super().setUp(enable_file_logging=False) + + # Wildcards tests + + def test_wc_ignore(self): # wildcards with ignore option + self.process( + PromptPair("__bad_wildcard__", "{option1|option2}"), + 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, + }, + self.grammar_content, + self.wildcards_obj, + self.extranetwork_maps_obj, + ), + ) + + def test_wc_remove(self): # wildcards with remove option + self.process( + PromptPair( + "[neg5] this is: __bad_wildcard__ a (([complex|simpleneg6|regular] testneg1)(test:2.0):1.5) \nBREAK, BREAK with [abcneg4:defneg2(neg3:1.6):5] ", + "normal quality, {option1|option2}", + ), + PromptPair( + "this is: a (([complex|simple|regular] test)(test:2):1.5)\nBREAK with [abc:def:5]", + "[neg5], ([|neg6|]:1.65), (neg1:1.65), [neg4::5], normal quality, [neg2(neg3:1.6):5]", + ), + ppp=PromptPostProcessor( + self.ppp_logger, + self.interrupt, + self.def_env_info, + { + **self.defopts, + "process_wildcards": False, + "if_wildcards": PromptPostProcessor.IFWILDCARDS_CHOICES.remove.value, + }, + self.grammar_content, + self.wildcards_obj, + self.extranetwork_maps_obj, + ), + ) + + def test_wc_warn(self): # wildcards with warn option + self.process( + PromptPair("__bad_wildcard__", "{option1|option2}"), + 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, + }, + self.grammar_content, + self.wildcards_obj, + self.extranetwork_maps_obj, + ), + ) + + def test_wc_stop(self): # wildcards with stop option + self.process( + PromptPair("__bad_wildcard__", "{option1|option2}"), + PromptPair( + PromptPostProcessor.WILDCARD_STOP.format("__bad_wildcard__") + "__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.stop.value, + }, + self.grammar_content, + self.wildcards_obj, + self.extranetwork_maps_obj, + ), + interrupted=True, + ) + + def test_wcinvar_warn(self): # wildcards in var with warn option + self.process( + PromptPair("${v=__bad_wildcard__}${v}", ""), + 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, + }, + self.grammar_content, + self.wildcards_obj, + self.extranetwork_maps_obj, + ), + ) + + def test_wc_invalid_name(self): + self.process( + PromptPair("the choices are: ___invalid__", ""), + PromptPair("the choices are: ___invalid__", ""), + ppp="nocup", + ) + + def test_wc_wildcard1a_text(self): # simple text wildcard + self.process( + PromptPair("the choices are: __text/wildcard1__", ""), + PromptPair("the choices are: choice2", ""), + ppp="nocup", + ) + + def test_wc_wildcard1a_json(self): # simple json wildcard + self.process( + PromptPair("the choices are: __json/wildcard1__", ""), + PromptPair("the choices are: choice2", ""), + ppp="nocup", + ) + + def test_wc_wildcard1a_yaml(self): # simple yaml wildcard + self.process( + PromptPair("the choices are: __yaml/wildcard1__", ""), + PromptPair("the choices are: choice2", ""), + ppp="nocup", + ) + + def test_wc_wildcard1b_text(self): # simple text wildcard with multiple choices + self.process( + PromptPair("the choices are: __2-$$text/wildcard1__", ""), + PromptPair("the choices are: choice3, choice1", ""), + ppp="nocup", + ) + + def test_wc_wildcard1b_json(self): # simple json wildcard with multiple choices + self.process( + PromptPair("the choices are: __2-$$json/wildcard1__", ""), + PromptPair("the choices are: choice3, choice1", ""), + ppp="nocup", + ) + + def test_wc_wildcard1b_yaml(self): # simple yaml wildcard with multiple choices + self.process( + PromptPair("the choices are: __2-$$yaml/wildcard1__", ""), + PromptPair("the choices are: choice3, choice1", ""), + ppp="nocup", + ) + + def test_wc_wildcard2_text(self): # simple text wildcard with default options + self.process( + PromptPair("the choices are: __text/wildcard2__", ""), + PromptPair("the choices are: choice3-choice1", ""), + ppp="nocup", + ) + + def test_wc_wildcard2_json(self): # simple json wildcard with default options + self.process( + PromptPair("the choices are: __json/wildcard2__", ""), + PromptPair("the choices are: choice3-choice1", ""), + ppp="nocup", + ) + + def test_wc_wildcard2_yaml(self): # simple yaml wildcard with default options + self.process( + PromptPair("the choices are: __yaml/wildcard2__", ""), + PromptPair("the choices are: choice3-choice1", ""), + ppp="nocup", + ) + + def test_wc_test2_yaml(self): # simple yaml wildcard + self.process( + PromptPair("the choice is: __testwc/test2__", ""), + PromptPair("the choice is: 2", ""), + ppp="nocup", + ) + + def test_wc_test3_yaml(self): # simple yaml wildcard + self.process( + PromptPair("the choice is: __testwc/test3__", ""), + PromptPair("the choice is: one choice", ""), + ppp="nocup", + ) + + def test_wc_wildcard_filter_index(self): # wildcard with positional index filter + self.process( + PromptPair("the choice is: __yaml/wildcard2'2'__", ""), + PromptPair("the choice is: choice3-choice3", ""), + ppp="nocup", + ) + + def test_wc_wildcard_filter_label(self): # wildcard with label filter + self.process( + PromptPair("the choice is: __yaml/wildcard2'label1'__", ""), + PromptPair("the choice is: choice3-choice1", ""), + ppp="nocup", + ) + + def test_wc_wildcard_filter_label2(self): # wildcard with label filter in multiple choices + self.process( + PromptPair("the choice is: __yaml/wildcard2'label2'__", ""), + PromptPair("the choice is: choice1-choice1", ""), + ppp="nocup", + ) + + def test_wc_wildcard_filter_label3(self): # wildcard with multiple label filter + self.process( + PromptPair("the choice is: __yaml/wildcard2'label1,label2'__", ""), + PromptPair("the choice is: choice3-choice1", ""), + ppp="nocup", + ) + + def test_wc_wildcard_filter_indexlabel(self): # wildcard with mixed index and label filter + self.process( + PromptPair("the choice is: __yaml/wildcard2'2,label2'__", ""), + PromptPair("the choice is: choice3-choice1", ""), + ppp="nocup", + ) + + def test_wc_wildcard_filter_compound(self): # wildcard with compound filter + self.process( + PromptPair("the choice is: __yaml/wildcard2'label1+label3'__", ""), + PromptPair("the choice is: choice3-choice3", ""), + ppp="nocup", + ) + + def test_wc_wildcard_filter_compound2(self): # wildcard with inherited compound filter + self.process( + PromptPair("the choice is: __yaml/wildcard2bis'#label1+label3'__", ""), + PromptPair("the choice is: choice3bis", ""), + ppp="nocup", + ) + + def test_wc_wildcard_filter_compound3(self): # wildcard with doubly inherited compound filter + self.process( + PromptPair("the choice is: __yaml/wildcard2bisbis'#label1+label3'__", ""), + PromptPair("the choice is: choice3bisbis", ""), + ppp="nocup", + ) + + def test_wc_wildcard_default_filter(self): # wildcard with default filter + self.process( + PromptPair( + "the choice is: __yaml/wildcard2__, __yaml/wildcard2__", + "", + ), + PromptPair("the choice is: choice3-choice3, choice3-choice1- choice2 ", ""), + ppp="nocup", + ) + + def test_wc_nested_wildcard_text(self): # nested text wildcard with repeating multiple choices + self.process( + PromptPair("the choices are: __r3$$-$$text/wildcard3__", ""), + PromptPair("the choices are: choice3,choice1- choice2 ,choice3", ""), + ppp="nocup", + ) + + def test_wc_nested_wildcard_json(self): # nested json wildcard with repeating multiple choices + self.process( + PromptPair("the choices are: __r3$$-$$json/wildcard3__", ""), + PromptPair("the choices are: choice3,choice1- choice2 ,choice3", ""), + ppp="nocup", + ) + + def test_wc_nested_wildcard_yaml(self): # nested yaml wildcard with repeating multiple choices + self.process( + PromptPair("the choices are: __r3$$-$$yaml/wildcard3__", ""), + PromptPair("the choices are: choice3,choice1- choice2 ,choice3", ""), + ppp="nocup", + ) + + def test_wc_wildcard_optional(self): # empty wildcard with no error + self.process( + PromptPair("the choices are: __yaml/empty_wildcard__", ""), + PromptPair("the choices are: ", ""), + ppp="nocup", + ) + + def test_wc_wildcard4_yaml(self): # simple yaml wildcard with one option + self.process( + PromptPair("the choices are: __yaml/wildcard4__", ""), + PromptPair("the choices are: inline text", ""), + ppp="nocup", + ) + + def test_wc_wildcard6_yaml(self): # simple yaml wildcard with object formatted choices + self.process( + PromptPair("the choices are: __yaml/wildcard6__", ""), + PromptPair("the choices are: choice2", ""), + ppp="nocup", + ) + + def test_wc_choice_wildcard_mix(self): # choices with wildcard mix + self.process( + PromptPair("the choices are: {__~2$$yaml/wildcard2__|choice0}", ""), + [ + PromptPair("the choices are: choice0", ""), + PromptPair("the choices are: choice1, choice3", ""), + PromptPair("the choices are: choice1, choice3", ""), + ], + ppp="nocup", + ) + + def test_wc_unsupportedsampler(self): # unsupported sampler + self.process( + PromptPair("the choices are: __@yaml/wildcard2__", ""), + PromptPair("", ""), + ppp="nocup", + interrupted=True, + ) + + def test_wc_wildcard_globbing(self): # wildcard with globbing + self.process( + PromptPair("the choices are: __yaml/wildcard[12]__, __yaml/wildcard?__", ""), + PromptPair("the choices are: choice3-choice2, - choice2 -choice3", ""), + ppp="nocup", + ) + + def test_wc_wildcardwithvar(self): # wildcard with inline variable + self.process( + PromptPair("the choices are: __yaml/wildcard5(var=test)__, __yaml/wildcard5__", ""), + PromptPair("the choices are: inline test, inline default", ""), + ppp="nocup", + ) + + def test_wc_wildcardPS_yaml(self): # yaml wildcard with object formatted choices and options and prefix and suffix + self.process( + PromptPair("the choices are: __yaml/wildcardPS__", ""), + PromptPair("the choices are: prefix-choice2/choice3-suffix", ""), + ppp="nocup", + ) + + def test_wc_anonymouswildcard_yaml(self): # yaml anonymous wildcard + self.process( + PromptPair("the choices are: __yaml/anonwildcards__", ""), + PromptPair("the choices are: six", ""), + ppp="nocup", + ) + + def test_wc_wildcard_input(self): # simple yaml wildcard input + self.process( + PromptPair("the choices are: __yaml_input/wildcardI__", ""), + PromptPair("the choices are: choice2", ""), + ppp="nocup", + ) + + def test_wc_circular(self): # wildcard circular reference + self.process( + PromptPair("the choices are: __yaml/circular1__", ""), + PromptPair("", ""), + ppp="nocup", + interrupted=True, + ) + + def test_wc_including(self): # wildcard including another wildcard + self.process( + PromptPair("the choices are: __yaml/including__", ""), + PromptPair("the choices are: choice4", ""), + ppp="nocup", + ) + + def test_wc_circular_including(self): # wildcard including another wildcard in a circular reference + self.process( + PromptPair("the choices are: __yaml/including1__", ""), + PromptPair("", ""), + ppp="nocup", + interrupted=True, + ) + + def test_wc_dynamicwildcard(self): # wildcard built from variables + self.process( + PromptPair( + "the choices are: ${x={1|2|3}}${w=yaml/wildcard${x}}__yaml/wildcard${x}__ __${w}__ ____", + "", + ), + PromptPair("the choices are: choice1-choice3-choice1 choice3- choice2 - choice2 choice3", ""), + ppp="nocup", + ) + + +if __name__ == "__main__": + unittest.main()