diff --git a/.gitignore b/.gitignore index 24a012d..225ca44 100644 --- a/.gitignore +++ b/.gitignore @@ -1,4 +1,5 @@ **/__pycache__ +.venv .vscode/**/* !.vscode/settings.json @@ -7,3 +8,5 @@ tests/tests_local.py tests/local_wildcards tests/logs + +scripts/last_prompts.txt diff --git a/.pylintrc b/.pylintrc index fd28a55..9e4b505 100644 --- a/.pylintrc +++ b/.pylintrc @@ -59,16 +59,6 @@ single-line-class-stmt=no single-line-if-stmt=no [IMPORTS] -allow-any-import-level= -allow-reexport-from-package=no -allow-wildcard-with-all=no -deprecated-modules= -ext-import-graph= -import-graph= -int-import-graph= -known-standard-library= -known-third-party=enchant -preferred-modules= [LOGGING] logging-format-style=new @@ -94,9 +84,7 @@ disable=raw-checker-failed, missing-module-docstring, missing-class-docstring, logging-fstring-interpolation, - import-outside-toplevel, consider-iterating-dictionary, - wrong-import-position, unnecessary-lambda, consider-using-dict-items, dangerous-default-value, diff --git a/.python-version b/.python-version new file mode 100644 index 0000000..5b609ee --- /dev/null +++ b/.python-version @@ -0,0 +1 @@ +3.10.11 diff --git a/.vscode/settings.json b/.vscode/settings.json index b8a6993..ad54e86 100644 --- a/.vscode/settings.json +++ b/.vscode/settings.json @@ -12,5 +12,12 @@ "python.analysis.typeCheckingMode": "off", "black-formatter.args": [ "--line-length=120" + ], + "python-envs.pythonProjects": [ + { + "path": ".", + "envManager": "ms-python.python:venv", + "packageManager": "ms-python.python:pip" + } ] } \ No newline at end of file diff --git a/README.md b/README.md index 3a00294..f76bd7f 100644 --- a/README.md +++ b/README.md @@ -17,6 +17,7 @@ These are some features: * Filter content based on the loaded SD model/variant or a variable. * Map extranetworks (LoRAs) depending on conditions (like the loaded model variant). This allows you to add "virtual" loras to the prompt that will be translated to the correct one. * Clean up the prompt of unnecessary separators or spaces. +* Combinatorial mode. Note: when used in an *A1111* compatible webui, the extension must be loaded after any other extension that modifies the prompt (like another wildcards extension). Usually extensions load by their folder name in alphanumeric order, so if the extensions are not loading in the correct order just rename this extension's folder so the ordering works out. When in doubt, just rename this extension's folder with a "z" in front (for example) so that it is the last one to load, or manually set such folder name when installing it. @@ -75,10 +76,14 @@ See the [syntax documentation](docs/SYNTAX.md). See the [cookbook](docs/COOKBOOK.md) for interesting usages. +## Contributing + +To develop, I suggest creating a virtual environment just for the extension, so the tests work and can be debugged properly. + ## License MIT ## Contact -If you have any questions or concerns, please start a thread in the discussions. +If you have any questions or concerns, please start a thread in the discussions. For bug reports and feature requests open an issue. diff --git a/docs/CONFIG.md b/docs/CONFIG.md index d52b4bd..342d27a 100644 --- a/docs/CONFIG.md +++ b/docs/CONFIG.md @@ -29,6 +29,8 @@ Inputs: * **process_wildcards**: Activates the wildcard processing. * **do_cleanup**: Activates the cleanup processing. * **cleanup_variables**: Do a cleanup of the output variables (depends on do_cleanup). +* **do_combinatorial**: Activates combinatorial mode, where the output are all the combinations of choices/wildcards of the prompt. +* **combinatorial_limit**: Limit for the number of generated combinations. * **wc_options**: Connection to a Wildcards options node. * **stn_options**: Connection to a Send-To-Negative options node. * **cup_options**: Connection to a Cleanup options node. @@ -42,6 +44,8 @@ Outputs: * **neg_prompt**: the resulting negative prompt * **variables**: the dictionary of variables set or echoed. +The outputs are lists, and in combinatorial mode there will be multiple elements that ComfyUI will process sequentially. + ### ACB PPP Select Variable node Lets you extract the variables used from the output (or just one of them). You can use this to send only part of the prompt to, for example, a detailer node. For example: @@ -127,6 +131,8 @@ Options for extranetworks mapping, in case you want to change them from the defa * **Unlink seed**: Uses the specified seed for the prompt generation instead of the one from the image. This seed is only used for wildcards and choices. * **Prompt seed**: The seed to use for the prompt generation. If -1 a random one will be used. * **Incremental seed**: When using a batch you can use this to set the rest of the prompt seeds with consecutive values. +* **Combinatorial mode**: Generate all possible prompt combinations (from choices and wildcards) and cycle through them to fill the batch. +* **Combinations limit**: Maximum number of combinations to generate (0 = no limit). When generating a batch the limit is automatically raised to at least the batch size. ### General settings diff --git a/ppp.code-workspace b/ppp.code-workspace new file mode 100644 index 0000000..077d428 --- /dev/null +++ b/ppp.code-workspace @@ -0,0 +1,22 @@ +{ + "folders": [ + { + "name": "sd-webui-prompt-postprocessor", + "path": "." + } + ], + "settings": { + "[python]": { + "editor.defaultFormatter": "ms-python.black-formatter" + }, + "yaml.schemaStore.enable": false, + "cSpell.enabledFiletypes": [ + "lark" + ], + "powershell.cwd": "sd-webui-prompt-postprocessor" + }, + "launch": { + "version": "0.2.0", + "configurations": [] + } +} \ No newline at end of file diff --git a/ppp.py b/ppp.py index 069e631..e632d6e 100644 --- a/ppp.py +++ b/ppp.py @@ -8,7 +8,7 @@ import lark import numpy as np import yaml -from pydantic import ValidationError # pylint: disable=import-error +from pydantic import ValidationError from ppp_classes import ( FindInFilenamePattern, HostConfig, @@ -21,7 +21,7 @@ from ppp_classes import ( PPPInterrupt, PPPState, PPPStateOptions, -) # pylint: disable=import-error +) from ppp_variables import VariableRepository from ppp_logging import DEBUG_LEVEL, log from ppp_tree import TreeProcessor @@ -83,6 +83,8 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in DEFAULT_CUP_MERGE_ATTENTION = defopt["cup_merge_attention"] DEFAULT_CUP_REMOVE_EXTRANETWORK_TAGS = defopt["cup_remove_extranetwork_tags"] DEFAULT_STRICT_OPERATORS = defopt["strict_operators"] + DEFAULT_DO_COMBINATORIAL = defopt["do_combinatorial"] + DEFAULT_COMBINATORIAL_LIMIT = defopt["combinatorial_limit"] WILDCARD_WARNING = '(WARNING TEXT "INVALID WILDCARD" IN BRIGHT RED:1.5)\nBREAK ' WILDCARD_STOP = "INVALID WILDCARD! {0}\nBREAK " UNPROCESSED_STOP = "UNPROCESSED CONSTRUCTS!\nBREAK " @@ -138,7 +140,7 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in if user_config_file == "": if self.env_info.get("app", "") == SUPPORTED_APPS.comfyui.value: try: - import folder_paths # pylint: disable=import-error # type: ignore + import folder_paths # type: ignore user_dir = folder_paths.get_user_directory() if user_dir and os.path.isdir(user_dir): @@ -772,7 +774,115 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in "simple parser without new constructs", ) - def __processprompts(self, rng, prompt, negative_prompt): + def __postprocess_result( + self, + result: tuple[str, list[tuple[str, bool]], tuple[dict[str, str | None], dict[str, str | None]]], + ) -> tuple[str, str, dict[str, str | None]]: + variables = {} + unified_prompt, rem_wildcards, (_, echoed_variables_snapshot) = result + + # Split the unified prompt back into prompt and negative prompt + split_parts = unified_prompt.split("\x1d", 1) + prompt = split_parts[0] + negative_prompt = split_parts[1] if len(split_parts) > 1 else "" + + # Clean up + prompt = self.__cleanup(prompt, 1) + negative_prompt = self.__cleanup(negative_prompt, -1) + + self.log(logging.INFO, f"Result prompt: {prompt}") + self.log(logging.INFO, f"Result negative_prompt: {negative_prompt}") + try: + # Get and clean variables + var_keys = sorted(echoed_variables_snapshot.keys()) + for k in var_keys: + ev = echoed_variables_snapshot.get(k) + variables[k] = self.__cleanup(ev, 0) if self.state.options.cup_cleanup_variables else ev + + self.log(logging.DEBUG, f"Result variables: {variables}") + + # Result checks + warnings = [] + + # Check for special character sequences that should not be in the result + compound_prompt = prompt + "\n" + negative_prompt + found_sequences = re.findall(r"::|\$\$|\$\{|[{}]", compound_prompt) + if found_sequences: + s = ", ".join(map(lambda x: '"' + x + '"', set(found_sequences))) + warnings.append(f"Probably invalid character sequences: {s}.") + # Check for correctly nested parentheses and brackets + stack = [] + prev_char = "" + for char in compound_prompt: + if prev_char != "\\": + if char in "([": # opening characters + stack.append(char) + elif char in ")]": # closing characters + if not stack: + warnings.append(f"Unmatched '{char}' character.") + break + last_open = stack.pop() + if (last_open == "(" and char != ")") or (last_open == "[" and char != "]"): + warnings.append(f"Mismatched '{last_open}' and '{char}' characters.") + break + prev_char = char + else: + prev_char = "" # reset prev_char to avoid treating escaped characters as escapes + if stack: + warnings.append(f"Unmatched '{''.join(stack)}' characters.") + if warnings: + self.log( + logging.WARNING, + "Found some weird things in the result. Something might be wrong!\n" + + "\n".join(f" - {w}" for w in warnings), + ) + + # Check for wildcards not processed + if rem_wildcards: + w_found_p = [wc for wc, n in rem_wildcards if not n] + w_found_n = [wc for wc, n in rem_wildcards if n] + if self.state.options.if_wildcards == IFWILDCARDS_CHOICES.stop: + self.log(logging.ERROR, "Found unprocessed wildcards!") + else: + self.log(logging.INFO, "Found unprocessed wildcards.") + ppwl = ", ".join(w_found_p) + npwl = ", ".join(w_found_n) + if ppwl: + self.log(logging.ERROR, f"In the prompt: {ppwl}") + if npwl: + self.log(logging.ERROR, f"In the negative prompt: {npwl}") + if self.state.options.if_wildcards == IFWILDCARDS_CHOICES.warn: + prompt = self.WILDCARD_WARNING + prompt + elif self.state.options.if_wildcards == IFWILDCARDS_CHOICES.stop: + raise PPPInterrupt( + "Found unprocessed wildcards!", + self.WILDCARD_STOP.format(ppwl) if ppwl else "", + self.WILDCARD_STOP.format(npwl) if npwl else "", + ) + + # Check for constructs not processed due to parsing problems + ppp_in_prompt = prompt.find("= 0 + ppp_in_negative_prompt = negative_prompt.find("= 0 + if ppp_in_prompt or ppp_in_negative_prompt: + raise PPPInterrupt( + "Found unprocessed constructs!", + self.UNPROCESSED_STOP if ppp_in_prompt else "", + self.UNPROCESSED_STOP if ppp_in_negative_prompt else "", + ) + except PPPInterrupt as e: + self.log(logging.ERROR, e.message) + if e.pos_prefix: + prompt = e.pos_prefix + prompt + if e.neg_prefix: + negative_prompt = e.neg_prefix + negative_prompt + self.log(logging.ERROR, "Interrupting!") + self.interrupt() + + v = self.state.variables.get_all_system() + v.update(variables) + return prompt, negative_prompt, v + + def __processprompts(self, rng, prompt, negative_prompt) -> list[tuple[str, str, dict[str, str | None]]]: """ Process the prompt and negative prompt. @@ -782,16 +892,15 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in negative_prompt (str): The negative prompt. Returns: - tuple: A tuple containing the processed prompt and negative prompt. + list: A list of tuples, each containing the processed prompt, negative prompt, and all variables. """ self.state.variables.clear_user() self.state.variables.clear_echoed() - all_variables = self.state.variables.get_all_system() # Parse both prompts processor = TreeProcessor(self.state, rng) - unified_prompt = prompt + "\x1D" + negative_prompt - (prompt_parser, parser_description) = self.__get_best_parser(unified_prompt) + unified_prompt = prompt + "\x1d" + negative_prompt + prompt_parser, parser_description = self.__get_best_parser(unified_prompt) self.log(logging.DEBUG, f"Using {parser_description} for prompt") parsed = parse_prompt( self.state, @@ -801,100 +910,28 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in ) # Process the unified prompt - unified_prompt, rem_wildcards = processor.start_visit(parsed) + t1 = time.monotonic_ns() + try: + results = processor.start_visit(parsed) + except PPPInterrupt as e: + self.log(logging.ERROR, e.message) + if e.pos_prefix: + prompt = e.pos_prefix + prompt + if e.neg_prefix: + negative_prompt = e.neg_prefix + negative_prompt + self.log(logging.ERROR, "Interrupting!") + self.interrupt() + t2 = time.monotonic_ns() + self.log(logging.INFO, f"Visit time: {(t2 - t1) / 1_000_000_000:.3f} seconds") - # Complete variables - var_keys = sorted(self.state.variables.all_user_or_echoed_keys()) - for k in var_keys: - ev = self.state.variables.get_echoed_value(k) - if ev is None: - ev = self.state.variables.get_user(k) - if ev is None or not isinstance(ev, str): - self.log(logging.DEBUG, f"Completing variable: {k}") - ev = processor.get_final_variable(k) - all_variables[k] = self.__cleanup(ev, 0) if self.state.options.cup_cleanup_variables else ev - self.log(logging.DEBUG, f"All variables: {all_variables}") - - # Split the unified prompt back into prompt and negative prompt - split_parts = unified_prompt.split("\x1D", 1) - prompt = split_parts[0] - negative_prompt = split_parts[1] if len(split_parts) > 1 else "" - - # Clean up - prompt = self.__cleanup(prompt, 1) - negative_prompt = self.__cleanup(negative_prompt, -1) - - # Result checks - warnings = [] - - # Check for special character sequences that should not be in the result - compound_prompt = prompt + "\n" + negative_prompt - found_sequences = re.findall(r"::|\$\$|\$\{|[{}]", compound_prompt) - if found_sequences: - warnings.append( - f"Probably invalid character sequences: {', '.join(map(lambda x: '"' + x + '"', set(found_sequences)))}." - ) - # Check for correctly nested parentheses and brackets - stack = [] - prev_char = "" - for char in compound_prompt: - if prev_char != "\\": - if char in "([": # opening characters - stack.append(char) - elif char in ")]": # closing characters - if not stack: - warnings.append(f"Unmatched '{char}' character.") - break - last_open = stack.pop() - if (last_open == "(" and char != ")") or (last_open == "[" and char != "]"): - warnings.append(f"Mismatched '{last_open}' and '{char}' characters.") - break - prev_char = char - else: - prev_char = "" # reset prev_char to avoid treating escaped characters as escapes - if stack: - warnings.append(f"Unmatched '{''.join(stack)}' characters.") - if warnings: - self.log( - logging.WARNING, - "Found some weird things in the result. Something might be wrong!\n" - + "\n".join(f" - {w}" for w in warnings), - ) - - # Check for wildcards not processed - if rem_wildcards: - w_found_p = [wc for wc, n in rem_wildcards if not n] - w_found_n = [wc for wc, n in rem_wildcards if n] - if self.state.options.if_wildcards == IFWILDCARDS_CHOICES.stop: - self.log(logging.ERROR, "Found unprocessed wildcards!") - else: - self.log(logging.INFO, "Found unprocessed wildcards.") - ppwl = ", ".join(w_found_p) - npwl = ", ".join(w_found_n) - if ppwl: - self.log(logging.ERROR, f"In the prompt: {ppwl}") - if npwl: - self.log(logging.ERROR, f"In the negative prompt: {npwl}") - if self.state.options.if_wildcards == IFWILDCARDS_CHOICES.warn: - prompt = self.WILDCARD_WARNING + prompt - elif self.state.options.if_wildcards == IFWILDCARDS_CHOICES.stop: - raise PPPInterrupt( - "Found unprocessed wildcards!", - self.WILDCARD_STOP.format(ppwl) if ppwl else "", - self.WILDCARD_STOP.format(npwl) if npwl else "", - ) - - # Check for constructs not processed due to parsing problems - ppp_in_prompt = prompt.find("= 0 - ppp_in_negative_prompt = negative_prompt.find("= 0 - if ppp_in_prompt or ppp_in_negative_prompt: - raise PPPInterrupt( - "Found unprocessed constructs!", - self.UNPROCESSED_STOP if ppp_in_prompt else "", - self.UNPROCESSED_STOP if ppp_in_negative_prompt else "", - ) - - return prompt, negative_prompt, all_variables + final_results = [] + for i, r in enumerate(results): + if self.state.options.do_combinatorial: + self.log(logging.INFO, f"Combination {i + 1}:") + final_results.append(self.__postprocess_result(r)) + if self.state.options.do_combinatorial: + self.log(logging.INFO, f"Total combinations: {len(final_results)}") + return final_results def process_prompt( self, @@ -913,7 +950,6 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in Returns: tuple: A tuple containing the processed prompt, negative prompt and all the prompt variables. """ - all_variables = {} try: if seed == -1: seed = np.random.randint(0, 2**32, dtype=np.int64) @@ -923,16 +959,13 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in self.log(logging.INFO, f"Input seed: {seed}") self.log(logging.INFO, f"Input prompt: {prompt}") self.log(logging.INFO, f"Input negative_prompt: {negative_prompt}") + self.log(logging.INFO, f"Combinatorial: {self.state.options.do_combinatorial}") t1 = time.monotonic_ns() - prompt, negative_prompt, all_variables = self.__processprompts( - np.random.default_rng(seed & 0xFFFFFFFF), prompt, negative_prompt - ) + results = self.__processprompts(np.random.default_rng(seed & 0xFFFFFFFF), prompt, negative_prompt) t2 = time.monotonic_ns() - self.log(logging.INFO, f"Result prompt: {prompt}") - self.log(logging.INFO, f"Result negative_prompt: {negative_prompt}") self.log(logging.INFO, f"Process prompt pair time: {(t2 - t1) / 1_000_000_000:.3f} seconds") # self.log(logging.DEBUG,f"Wildcards memory usage: {self.state.wildcards_obj.__sizeof__()}") - return prompt, negative_prompt, all_variables + return results except PPPInterrupt as e: self.log(logging.ERROR, e.message) if e.pos_prefix: @@ -941,7 +974,7 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in negative_prompt = e.neg_prefix + negative_prompt self.log(logging.ERROR, "Interrupting!") self.interrupt() - return prompt, negative_prompt, all_variables + return [prompt, negative_prompt, {}] except Exception: # pylint: disable=broad-exception-caught self.log(logging.ERROR, "Unexpected error", exc_info=True) - return original_prompt, original_negative_prompt, all_variables + return [original_prompt, original_negative_prompt, {}] diff --git a/ppp_classes.py b/ppp_classes.py index 1c17b77..d8905e0 100644 --- a/ppp_classes.py +++ b/ppp_classes.py @@ -201,7 +201,24 @@ class PPPStateOptions: cup_merge_attention: bool = True cup_remove_extranetwork_tags: bool = False strict_operators: bool = True + do_combinatorial: bool = False + combinatorial_limit: int = 100 # 0 = no limit + def __post_init__(self): + if not self.cup_do_cleanup: + object.__setattr__(self, "cup_cleanup_variables", False) + object.__setattr__(self, "cup_extra_spaces", False) + object.__setattr__(self, "cup_empty_constructs", False) + object.__setattr__(self, "cup_extra_separators", False) + object.__setattr__(self, "cup_extra_separators2", False) + object.__setattr__(self, "cup_extra_separators_include_eol", False) + object.__setattr__(self, "cup_breaks", False) + object.__setattr__(self, "cup_breaks_eol", False) + object.__setattr__(self, "cup_ands", False) + object.__setattr__(self, "cup_ands_eol", False) + object.__setattr__(self, "cup_extranetwork_tags", False) + object.__setattr__(self, "cup_merge_attention", False) + object.__setattr__(self, "cup_remove_extranetwork_tags", False) @dataclass(frozen=True) class PPPState: diff --git a/ppp_comfyui.py b/ppp_comfyui.py index 29edaef..1fb7571 100644 --- a/ppp_comfyui.py +++ b/ppp_comfyui.py @@ -1,8 +1,8 @@ import logging import os -import folder_paths # pylint: disable=import-error # type: ignore -import nodes # pylint: disable=import-error # type: ignore +import folder_paths # type: ignore +import nodes # type: ignore from ppp import PromptPostProcessor from ppp_classes import IFWILDCARDS_CHOICES, ONWARNING_CHOICES, SUPPORTED_APPS, PPPStateOptions @@ -180,6 +180,22 @@ class PromptPostProcessorComfyUINode: "label_off": "No", }, ), + "do_combinatorial": ( + "BOOLEAN", + { + "default": PromptPostProcessor.DEFAULT_DO_COMBINATORIAL, + "tooltip": "Enable combinatorial mode", + "label_on": "Yes", + "label_off": "No", + }, + ), + "combinatorial_limit": ( + "INT", + { + "default": PromptPostProcessor.DEFAULT_COMBINATORIAL_LIMIT, + "tooltip": "Limit for combinatorial mode", + }, + ), "wc_options": ( "PPP_OPTIONS_WC", { @@ -229,6 +245,11 @@ class PromptPostProcessorComfyUINode: "STRING", "PPP_DICT", ) + OUTPUT_IS_LIST = ( + True, + True, + True, + ) RETURN_NAMES = ( "pos_prompt", "neg_prompt", @@ -255,6 +276,8 @@ class PromptPostProcessorComfyUINode: process_wildcards, do_cleanup, cleanup_variables, + do_combinatorial, + combinatorial_limit, wc_options=None, stn_options=None, cup_options=None, @@ -292,7 +315,9 @@ class PromptPostProcessorComfyUINode: options = PPPStateOptions( debug_level=DEBUG_LEVEL(debug_level), on_warning=ONWARNING_CHOICES(on_warnings) if on_warnings else PromptPostProcessor.DEFAULT_ON_WARNING, - strict_operators=strict_operators if strict_operators is not None else PromptPostProcessor.DEFAULT_STRICT_OPERATORS, + strict_operators=( + strict_operators if strict_operators is not None else PromptPostProcessor.DEFAULT_STRICT_OPERATORS + ), process_wildcards=process_wildcards, if_wildcards=(wc_options["wc_if_wildcards"] if wc_options else IFWILDCARDS_CHOICES.stop.value), choice_separator=( @@ -345,6 +370,8 @@ class PromptPostProcessorComfyUINode: if cup_options else PromptPostProcessor.DEFAULT_CUP_REMOVE_EXTRANETWORK_TAGS ), + do_combinatorial=do_combinatorial, + combinatorial_limit=combinatorial_limit, ) self.wildcards_obj.refresh_wildcards( options.debug_level, @@ -365,12 +392,14 @@ class PromptPostProcessorComfyUINode: self.wildcards_obj, self.extranetwork_mappings_obj, ) - pos_prompt, neg_prompt, variables = ppp.process_prompt(pos_prompt, neg_prompt, seed if seed is not None else 1) - return ( - pos_prompt, - neg_prompt, - variables, - ) + results = ppp.process_prompt(pos_prompt, neg_prompt, seed if seed is not None else 1) + # pos_prompt, neg_prompt, variables = results[0] + # return ( + # pos_prompt, + # neg_prompt, + # variables, + # ) + return tuple(zip(*results)) # unzip the list of tuples into tuple of lists def interrupt(self): nodes.interrupt_processing(True) @@ -917,8 +946,8 @@ class PromptPostProcessorWildcardConcatComfyUINode: try: - from server import PromptServer # type: ignore # pylint: disable=import-error - from aiohttp import web as _aiohttp_web # type: ignore # pylint: disable=import-error + from server import PromptServer # type: ignore + from aiohttp import web as _aiohttp_web # type: ignore @PromptServer.instance.routes.get("/acb_ppp/wildcards") async def _acb_ppp_get_wildcards(request): diff --git a/ppp_enmappings.py b/ppp_enmappings.py index e69b219..8bdedfa 100644 --- a/ppp_enmappings.py +++ b/ppp_enmappings.py @@ -1,4 +1,5 @@ import os +from pathlib import Path from typing import Optional import logging import yaml @@ -102,7 +103,7 @@ class PPPExtraNetworkMappings: if fullpath != self.LOCALINPUT_FILENAME: path = os.path.dirname(fullpath) if not os.path.exists(fullpath) or not any( - os.path.commonpath([path, folder]) == folder for folder in self.__enmappings_folders + Path(path).is_relative_to(folder) for folder in self.__enmappings_folders ): self.__remove_extranetwork_mappings_from_path(fullpath) elif enmappings_input is None: diff --git a/ppp_tree.py b/ppp_tree.py index 299b3b5..530d396 100644 --- a/ppp_tree.py +++ b/ppp_tree.py @@ -1,10 +1,11 @@ from collections import namedtuple +from itertools import combinations, combinations_with_replacement, permutations, product import logging import math import re import textwrap import time -from typing import Optional +from typing import Any, Optional import lark import numpy as np @@ -31,7 +32,7 @@ class TreeProcessor(lark.visitors.Interpreter): result (str): The final processed prompt. """ - NEGATIVE_SEP = "\x1D" + NEGATIVE_SEP = "\x1d" AccumulatedShell = namedtuple("AccumulatedShell", ["type", "data"]) NegTag = namedtuple("NegTag", ["start", "end", "content", "parameters", "shell"]) @@ -46,10 +47,12 @@ class TreeProcessor(lark.visitors.Interpreter): self.__is_negative = False self.__wildcard_filters = {} self.__seen_wildcards: list[str] = [] - self.add_at: dict[str, list] = {"start": [], "insertion_point": [[] for _ in range(10)], "end": []} - self.insertion_at: list[tuple[int, int]] = [None for _ in range(10)] - self.detectedWildcards: list[tuple[str,bool]] = [] - self.result = "" + self.__add_at: dict[str, list] = {"start": [], "insertion_point": [[] for _ in range(10)], "end": []} + self.__insertion_at: list[tuple[int, int]] = [None for _ in range(10)] + self.__detectedWildcards: list[tuple[str, bool]] = [] + self.__result = "" + self.__comb_forced_path: list[int] = [] + self.__comb_trace: list[int] = [] def log(self, kind, message: str, min_level: DEBUG_LEVEL | None = None): log(self.state.logger, self.state.options.debug_level, kind, message, min_level) @@ -57,10 +60,25 @@ class TreeProcessor(lark.visitors.Interpreter): def warn_or_stop(self, message: str, e: Exception = None): warn_or_stop(self.state, self.__is_negative, message, e) + def __reset_run_state(self): + """Reset all per-run mutable state for a fresh combinatorial pass.""" + self.__shell = [] + self.__negtags = [] + self.__already_processed = [] + self.__is_negative = False + self.__wildcard_filters = {} + self.__seen_wildcards = [] + self.__add_at = {"start": [], "insertion_point": [[] for _ in range(10)], "end": []} + self.__insertion_at = [None for _ in range(10)] + self.__detectedWildcards = [] + self.__result = "" + if self.state.extranetwork_mappings_obj is not None: + self.state.extranetwork_mappings_obj.cached_mappings.clear() + def start_visit( self, parsed: lark.Tree, - ) -> tuple[str, list[tuple[str,bool]]]: + ) -> list[tuple[str, list[tuple[str, bool]], tuple[dict[str, Any], dict[str, str]]]]: """ Process the positive and negative prompts in a unified way using the same processor. STN insertions are applied to the negative result directly inside this processor. @@ -69,17 +87,78 @@ class TreeProcessor(lark.visitors.Interpreter): parsed (Tree): The parsed unified prompt. Returns: - tuple[str, list[tuple[str,bool]]]: The processed prompt and its detected wildcards. + list[tuple[str, list[tuple[str,bool]], tuple[dict[str, Any], dict[str, str]]]]: A list of + (processed prompt, detected wildcards, variables snapshot) triples — one entry per + combination in combinatorial mode, or a single entry otherwise. The variables snapshot + is the return value of ``state.variables.backup_user_and_echoed()`` captured after processing. """ - t1 = time.monotonic_ns() - self.log(logging.INFO, f"Processing prompt...") - self.detectedWildcards = [] + self.log(logging.INFO, "Processing prompt...") + + self.__detectedWildcards = [] self.__is_negative = False - self.result = "" - self.visit(parsed) - t2 = time.monotonic_ns() - self.log(logging.INFO, f"Process prompt time: {(t2 - t1) / 1_000_000_000:.3f} seconds") - return self.result, self.detectedWildcards + self.__result = "" + + if not self.state.options.do_combinatorial: + self.visit(parsed) + self.__finalize_echoed_variables() + return [(self.__result, self.__detectedWildcards, self.state.variables.backup_user_and_echoed())] + + # Combinatorial mode: explore every possible path through choices and wildcards via DFS. + # __comb_forced_path drives which option is selected at each decision point; + # __comb_trace records how many options were available at each point so the DFS can + # correctly enumerate unexplored branches after each run. + initial_vars = self.state.variables.backup_user_and_echoed() + results: list[tuple[str, list[tuple[str, bool]], tuple]] = [] + limit = self.state.options.combinatorial_limit + + def _run(forced_path: tuple[int, ...]) -> tuple[int, ...]: + self.log(logging.DEBUG, f"Running combinatorial path: {forced_path}") + self.__comb_forced_path = list(forced_path) + self.__comb_trace = [] + self.__reset_run_state() + self.state.variables.restore_user_and_echoed(initial_vars) + self.visit(parsed) + self.__finalize_echoed_variables() + results.append( + (self.__result, self.__detectedWildcards.copy(), self.state.variables.backup_user_and_echoed()) + ) + return tuple(self.__comb_trace) + + def _dfs(forced_path: tuple[int, ...]): + if 0 < limit <= len(results): + return + trace = _run(forced_path) + # For each decision that was reached but not forced, spawn branches for all + # options beyond the default (index 0). + # Iterate in reverse so later decisions vary fastest, producing lexicographic order. + for i in range(len(trace) - 1, len(forced_path) - 1, -1): + if 0 < limit <= len(results): + return + num_options = trace[i] + for opt in range(1, num_options): + if 0 < limit <= len(results): + return + # Pad with zeros for intermediate decisions so they keep the default. + new_path = forced_path + (0,) * (i - len(forced_path)) + (opt,) + _dfs(new_path) + + _dfs(()) + if 0 < limit <= len(results): + self.log( + logging.WARNING, f"Combinatorial limit of {limit} reached; some combinations may have been skipped." + ) + return results + + def __finalize_echoed_variables(self): + var_keys = self.state.variables.all_user_or_echoed_keys() + for k in var_keys: + ev = self.state.variables.get_echoed_value(k) + if ev is None: + ev = self.state.variables.get_user(k) + if ev is None or ev.__class__ != str: # strict check to avoid problems with Tokens + self.log(logging.DEBUG, f"Completing variable: {k}") + ev = self.get_final_variable(k) + self.state.variables.echo(k, ev) # ensure all variables are echoed so they are included in the snapshot def __visit( self, @@ -98,16 +177,16 @@ class TreeProcessor(lark.visitors.Interpreter): Returns: str: The result of the visit. """ - backup_result = self.result + backup_result = self.__result # self.log(logging.DEBUG, f"Visiting node {node}.") if restore_state: # self.log(logging.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_add_at = self.__add_at.copy() + backup_insertion_at = self.__insertion_at.copy() + backup_detectedwildcards = self.__detectedWildcards.copy() backup_vars = self.state.variables.backup_user_and_echoed() if node is not None: if isinstance(node, list): @@ -116,22 +195,22 @@ class TreeProcessor(lark.visitors.Interpreter): elif isinstance(node, lark.Tree): self.visit(node) elif isinstance(node, lark.Token): - self.result += node + self.__result += node len_backup = len(backup_result) # if self.result[:len_backup] == backup_result: # this is only necessary if we call parse_prompt with a parser from "start", because it resets the result - added_result = self.result[len_backup:] + added_result = self.__result[len_backup:] # else: # added_result = self.result if discard_content or restore_state: - self.result = backup_result + self.__result = backup_result if restore_state: # self.log(logging.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.__add_at = backup_add_at + self.__insertion_at = backup_insertion_at + self.__detectedWildcards = backup_detectedwildcards self.state.variables.restore_user_and_echoed(backup_vars) return added_result @@ -162,7 +241,7 @@ class TreeProcessor(lark.visitors.Interpreter): return None, None, None if specifier == "#": # special value to indicate length of the array variable return None, None, True - if specifier.startswith("&"): # special value to indicate a separator + if specifier.startswith("&"): # special value to indicate a separator return None, specifier[2:-1], False if not specifier.isdecimal(): # bare identifier: resolve as variable @@ -199,7 +278,7 @@ class TreeProcessor(lark.visitors.Interpreter): elif isinstance(v, lark.Token): v = str(v) if visit and not visited: - self.result += v + self.__result += v return v v = self.state.variables.get(name) @@ -212,7 +291,7 @@ class TreeProcessor(lark.visitors.Interpreter): if cnt: v = len(v) if visit: - self.result += str(v) + self.__result += str(v) elif idx is not None: if 0 <= idx < len(v): v = visit_value(v[idx]) @@ -225,7 +304,7 @@ class TreeProcessor(lark.visitors.Interpreter): for i, item in enumerate(v): v2.append(visit_value(item)) if visit and i < len(v) - 1: - self.result += sep + self.__result += sep v = v2 else: v = None # error @@ -283,7 +362,7 @@ class TreeProcessor(lark.visitors.Interpreter): """ if self.__debug_level == DEBUG_LEVEL.full: info = f"({info}) " if info is not None and info != "" else "" - output = self.result[len(start_result) :] + output = self.__result[len(start_result) :] if output != "": output = f" >> '{escape_single_quotes(output)}'" self.log(logging.DEBUG, f"TreeProcessor.{construct} {info}({duration / 1_000_000_000:.3f} seconds){output}") @@ -333,7 +412,9 @@ class TreeProcessor(lark.visitors.Interpreter): (bool, list), ] if not any(isinstance(operand1, t1) and isinstance(operand2, t2) for t1, t2 in compatible_types): - self.warn_or_stop(f"Mixed type values ({type(operand1).__name__}, {type(operand2).__name__}) used in comparison: '{escape_single_quotes(desc)}'") + self.warn_or_stop( + f"Mixed type values ({type(operand1).__name__}, {type(operand2).__name__}) used in comparison: '{escape_single_quotes(desc)}'" + ) return False return operation(operand1, operand2) @@ -696,10 +777,10 @@ class TreeProcessor(lark.visitors.Interpreter): """ Process a negative prompt separator in the tree. """ - start_result = self.result + start_result = self.__result t1 = time.monotonic_ns() x = tree.children[0] - self.result += x.value + self.__result += x.value self.__is_negative = True t2 = time.monotonic_ns() self.__debug_end("negative_sep", start_result, t2 - t1) @@ -708,7 +789,7 @@ class TreeProcessor(lark.visitors.Interpreter): """ Process a prompt composition construct in the tree. """ - start_result = self.result + start_result = self.__result t1 = time.monotonic_ns() self.__visit(tree.children[0]) and_processing = self.state.host_config.and_ @@ -719,11 +800,11 @@ class TreeProcessor(lark.visitors.Interpreter): "remove": ("removed", " "), } if tree.children[1] is not None: - self.result += f":{tree.children[1]}" + self.__result += f":{tree.children[1]}" for i in range(2, len(tree.children), 3): if and_processing in and_replacements.keys(): - self.result = ( - self.result.rstrip() + self.__result = ( + self.__result.rstrip() + and_replacements[and_processing][1] + self.__visit(tree.children[i + 1], False, True).lstrip() ) @@ -732,18 +813,20 @@ class TreeProcessor(lark.visitors.Interpreter): self.warn_or_stop("AND constructs are not allowed!") else: # and_processing == "ok": if self.state.options.cup_ands: - self.result = re.sub(r"[, ]+$", "\n" if self.state.options.cup_ands_eol else " ", self.result) - if self.result[-1:].isalnum(): # add space if needed - self.result += " " - self.result += "AND" + self.__result = re.sub( + r"[, ]+$", "\n" if self.state.options.cup_ands_eol else " ", self.__result + ) + if self.__result[-1:].isalnum(): # add space if needed + self.__result += " " + self.__result += "AND" added_result = self.__visit(tree.children[i + 1], False, True) if self.state.options.cup_ands: added_result = re.sub(r"^[, ]+", " ", added_result) if added_result[0:1].isalnum(): # add space if needed added_result = " " + added_result - self.result += added_result + self.__result += added_result if tree.children[i + 2] is not None: - self.result += f":{tree.children[i+2]}" + self.__result += f":{tree.children[i+2]}" t2 = time.monotonic_ns() self.__debug_end("promptcomp", start_result, t2 - t1) @@ -751,7 +834,7 @@ class TreeProcessor(lark.visitors.Interpreter): """ Process a scheduling construct in the tree and add it to the accumulated shell. """ - start_result = self.result + start_result = self.__result t1 = time.monotonic_ns() before = tree.children[0] after = tree.children[-2] @@ -780,7 +863,7 @@ class TreeProcessor(lark.visitors.Interpreter): self.warn_or_stop("Scheduling constructs are not allowed!") else: # scheduling_processing == "ok" # self.__shell.append(TreeProcessor.AccumulatedShell("sc", pos)) - self.result += "[" + self.__result += "[" if before is not None: self.log(logging.DEBUG, f"Shell scheduled before with position {pos}") self.__shell.append(TreeProcessor.AccumulatedShell("scb", pos)) @@ -788,15 +871,15 @@ class TreeProcessor(lark.visitors.Interpreter): self.__shell.pop() self.log(logging.DEBUG, f"Shell scheduled after with position {pos}") self.__shell.append(TreeProcessor.AccumulatedShell("sca", pos)) - self.result += ":" + self.__result += ":" self.__visit(after) self.__shell.pop() if self.state.options.cup_empty_constructs and re.fullmatch( - re.escape(start_result) + r"\[:\s*", self.result + re.escape(start_result) + r"\[:\s*", self.__result ): - self.result = start_result + self.__result = start_result else: - self.result += f":{pos_str}]" + self.__result += f":{pos_str}]" # self.__shell.pop() t2 = time.monotonic_ns() self.__debug_end("scheduled", start_result, t2 - t1, pos_str) @@ -805,7 +888,7 @@ class TreeProcessor(lark.visitors.Interpreter): """ Process an alternation construct in the tree and add it to the accumulated shell. """ - start_result = self.result + start_result = self.__result t1 = time.monotonic_ns() alternation_processing = self.state.host_config.alternation if alternation_processing == "first": @@ -817,19 +900,19 @@ class TreeProcessor(lark.visitors.Interpreter): self.warn_or_stop("Alternation constructs are not allowed!") else: # alternation_processing == "ok" # self.__shell.append(TreeProcessor.AccumulatedShell("al", len(tree.children))) - self.result += "[" + self.__result += "[" for i, opt in enumerate(tree.children): self.log(logging.DEBUG, f"Shell alternate option {i+1}") self.__shell.append(TreeProcessor.AccumulatedShell("alo", {"pos": i + 1, "len": len(tree.children)})) if i > 0: - self.result += "|" + self.__result += "|" self.__visit(opt) self.__shell.pop() - self.result += "]" + self.__result += "]" if self.state.options.cup_empty_constructs and re.fullmatch( - re.escape(start_result) + r"\[\s*\]", self.result + re.escape(start_result) + r"\[\s*\]", self.__result ): - self.result = start_result + self.__result = start_result # self.__shell.pop() t2 = time.monotonic_ns() self.__debug_end("alternate", start_result, t2 - t1) @@ -838,7 +921,7 @@ class TreeProcessor(lark.visitors.Interpreter): """ Process a attention change construct in the tree and add it to the accumulated shell. """ - start_result = self.result + start_result = self.__result t1 = time.monotonic_ns() # weight_kind: -1: remove, 0=none, 1=decrease, 2=increase, 3=specific if len(tree.children) == 2: @@ -902,25 +985,25 @@ class TreeProcessor(lark.visitors.Interpreter): self.__shell.append(TreeProcessor.AccumulatedShell("at", (weight_kind, weight_str))) if weight_kind == 1: starttag = "[" - self.result += starttag + self.__result += starttag self.__visit(current_tree) endtag = "]" elif weight_kind == 2: starttag = "(" - self.result += starttag + self.__result += starttag self.__visit(current_tree) endtag = ")" else: # weight_kind == 3 starttag = "(" - self.result += starttag + self.__result += starttag self.__visit(current_tree) endtag = f":{weight_str})" if self.state.options.cup_empty_constructs and re.fullmatch( - re.escape(start_result + starttag) + r"\s*", self.result + re.escape(start_result + starttag) + r"\s*", self.__result ): - self.result = start_result + self.__result = start_result else: - self.result += endtag + self.__result += endtag self.__shell.pop() t2 = time.monotonic_ns() self.__debug_end("attention", start_result, t2 - t1, weight_str) @@ -929,7 +1012,7 @@ class TreeProcessor(lark.visitors.Interpreter): """ Process a send to negative command in the tree and add it to the list of negative tags. """ - start_result = self.result + start_result = self.__result info = None t1 = time.monotonic_ns() if not self.__is_negative: @@ -940,7 +1023,7 @@ class TreeProcessor(lark.visitors.Interpreter): parameters = "" content = self.__visit(tree.children[1::], False, True) self.__negtags.append( - TreeProcessor.NegTag(len(self.result), len(self.result), content, parameters, self.__shell.copy()) + TreeProcessor.NegTag(len(self.__result), len(self.__result), content, parameters, self.__shell.copy()) ) info = f"with {escape_single_quotes(parameters) or 'no parameters'} : {escape_single_quotes(content)}" else: @@ -953,7 +1036,7 @@ class TreeProcessor(lark.visitors.Interpreter): """ Process a send to negative insertion point command in the tree and add it to the list of negative tags. """ - start_result = self.result + start_result = self.__result info = None t1 = time.monotonic_ns() if self.__is_negative: @@ -962,7 +1045,9 @@ class TreeProcessor(lark.visitors.Interpreter): parameters = str(negtagparameters) else: parameters = "" - self.__negtags.append(TreeProcessor.NegTag(len(self.result), len(self.result), "", parameters, self.__shell.copy())) + self.__negtags.append( + TreeProcessor.NegTag(len(self.__result), len(self.__result), "", parameters, self.__shell.copy()) + ) info = f"with {parameters or 'no parameters'}" else: self.warn_or_stop("Ignored negative insertion point command in positive prompt") @@ -981,7 +1066,7 @@ class TreeProcessor(lark.visitors.Interpreter): Process a generic set command in the tree. """ t1 = time.monotonic_ns() - start_result = self.result + start_result = self.__result if self.state.variables.name_is_system(variable_name): self.warn_or_stop( f"Invalid variable name '{escape_single_quotes(variable_name)}' detected! System variables cannot be set." @@ -1069,7 +1154,9 @@ class TreeProcessor(lark.visitors.Interpreter): if is_starred: if access_full_array and isinstance(newvalue.children[0], lark.Tree): if newvalue.children[0].data == "vardescriptor_get": - vardescriptor_name, vardescriptor_specifier = self.__separate_vardescriptor(newvalue.children[0]) + vardescriptor_name, vardescriptor_specifier = self.__separate_vardescriptor( + newvalue.children[0] + ) if vardescriptor_specifier is not None: newvalue = None else: @@ -1079,9 +1166,9 @@ class TreeProcessor(lark.visitors.Interpreter): self.__resolve_operand(c) for c in self.__get_cond_operand(newvalue.children[0]) ) elif newvalue.children[0].data == "wildcard": - backup_result = self.result + backup_result = self.__result newvalue = self.__process_wildcard(newvalue.children[0]) - self.result = backup_result + self.__result = backup_result else: newvalue = None else: @@ -1150,7 +1237,7 @@ class TreeProcessor(lark.visitors.Interpreter): Process a generic echo command in the tree. """ t1 = time.monotonic_ns() - start_result = self.result + start_result = self.__result default_value = None # if default is not None: # default_value = self.__visit(default, True) # for log @@ -1161,7 +1248,7 @@ class TreeProcessor(lark.visitors.Interpreter): if default is not None: self.log(logging.DEBUG, f"Variable '{escape_single_quotes(vname)}' not found, using default value") value = self.__visit(default, False, True) - self.result += value + self.__result += value default_value = value else: self.warn_or_stop(f"Unknown variable {escape_single_quotes(vname)}") @@ -1206,7 +1293,7 @@ class TreeProcessor(lark.visitors.Interpreter): Process an if command in the tree. """ t1 = time.monotonic_ns() - start_result = self.result + start_result = self.__result for i, n in enumerate(tree.children): content = n.children[-1] if len(n.children) == 2: # its not an else @@ -1229,7 +1316,7 @@ class TreeProcessor(lark.visitors.Interpreter): Process an extranetwork command in the tree. """ t1 = time.monotonic_ns() - start_result = self.result + start_result = self.__result extnet = "(ignored)" if not self.state.options.cup_remove_extranetwork_tags: extnet_type: str = (tree.children[0].children[0] or "") + str(tree.children[0].children[1]) @@ -1290,12 +1377,23 @@ class TreeProcessor(lark.visitors.Interpreter): else: else_mapping = v if found_mappings: - found = found_mappings[ - self.__rng.choice( - len(found_mappings), - p=[v.weight or 1 for v in found_mappings], + if self.state.options.do_combinatorial: + N = len(found_mappings) + decision_idx = len(self.__comb_trace) + self.__comb_trace.append(N) + chosen_idx = ( + min(self.__comb_forced_path[decision_idx], N - 1) + if decision_idx < len(self.__comb_forced_path) + else 0 ) - ] + found = found_mappings[chosen_idx] + else: + found = found_mappings[ + self.__rng.choice( + len(found_mappings), + p=[v.weight or 1 for v in found_mappings], + ) + ] else: found = else_mapping self.state.extranetwork_mappings_obj.cached_mappings[extnet_id] = found @@ -1348,23 +1446,23 @@ class TreeProcessor(lark.visitors.Interpreter): 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 + self.__result += extnet elif triggers or compiled_extra_triggers: extnet = "(only triggers)" if triggers or compiled_extra_triggers: if extnet_id: if not self.state.options.cup_extranetwork_tags: - self.result += " " + self.__result += " " else: - self.result += ", " + self.__result += ", " if triggers: - self.result += self.__visit(triggers, True, True) + self.__result += self.__visit(triggers, True, True) if compiled_extra_triggers: if triggers: - self.result += ", " - self.result += self.__visit(compiled_extra_triggers, True, True) + self.__result += ", " + self.__result += self.__visit(compiled_extra_triggers, True, True) if triggers or compiled_extra_triggers: - self.result += ", " + self.__result += ", " t2 = time.monotonic_ns() self.__debug_end("commandext", start_result, t2 - t1, extnet) @@ -1373,7 +1471,7 @@ class TreeProcessor(lark.visitors.Interpreter): Process a setwcdeffilter (Set Wildcard Default Filter) command in the tree. """ t1 = time.monotonic_ns() - start_result = self.result + start_result = self.__result wildcard_key: str = self.__visit(tree.children[0].children[1], False, True) selected_wildcards = [x.key for x in self.state.wildcards_obj.get_wildcards(wildcard_key)] if not selected_wildcards: @@ -1397,11 +1495,11 @@ class TreeProcessor(lark.visitors.Interpreter): Process an extra network construct in the tree. """ t1 = time.monotonic_ns() - start_result = self.result + start_result = self.__result if not self.state.options.cup_remove_extranetwork_tags: - self.result += f"<{tree.children[0]}" + self.__result += f"<{tree.children[0]}" self.__visit(tree.children[1]) - self.result += ">" + self.__result += ">" t2 = time.monotonic_ns() self.__debug_end("extranetworktag", start_result, t2 - t1) @@ -1450,7 +1548,7 @@ class TreeProcessor(lark.visitors.Interpreter): for i, c in enumerate(filtered_choice_values): if c.get("command", False): content_text = self.__visit(c.get("content", ""), False, True).strip() - (cmd, cmd_args) = content_text.split() + cmd, cmd_args = content_text.split() if cmd == "include": wcs = self.state.wildcards_obj.get_wildcards(cmd_args) if not wcs: @@ -1467,7 +1565,7 @@ class TreeProcessor(lark.visitors.Interpreter): self.__seen_wildcards.append(wc.key) self.log(logging.DEBUG, f"Seen wildcard '{escape_single_quotes(wc.key)}'") self.log(logging.DEBUG, f"Including choices from wildcard '{escape_single_quotes(wc.key)}'") - (_, choice_values) = self.__check_wildcard_initialization(wc) + _, choice_values = self.__check_wildcard_initialization(wc) if choice_values is not None: ch_values = self.__get_choices_internal_get(choice_values, None, wc.key) for cv in ch_values: @@ -1550,9 +1648,40 @@ class TreeProcessor(lark.visitors.Interpreter): to_value = 1 elif (to_value > len(available_choices) and not repeating) or from_value > to_value: to_value = len(available_choices) - num_choices = ( - self.__rng.integers(from_value, to_value, endpoint=True) if from_value < to_value else from_value - ) + comb_chosen_selection: Optional[list[dict]] = None + if self.state.options.do_combinatorial: + # Enumerate every distinct selection of choices, accounting for count range and repetition. + all_selections: list[tuple] = [] + # When keep_choices_order is False the output depends on the selection order, + # so we must enumerate ordered sequences (permutations / product). + # When keep_choices_order is True selections are sorted afterward, so all + # orderings of the same items produce identical output and we only need + # unordered iterators (combinations / combinations_with_replacement). + for k in range(from_value, to_value + 1): + if repeating: + if self.state.options.keep_choices_order: + all_selections.extend(combinations_with_replacement(available_choices, k)) + else: + all_selections.extend(product(available_choices, repeat=k)) + else: + if self.state.options.keep_choices_order: + all_selections.extend(combinations(available_choices, k)) + else: + all_selections.extend(permutations(available_choices, k)) + num_selections = len(all_selections) + decision_idx = len(self.__comb_trace) + self.__comb_trace.append(num_selections) + chosen_idx = ( + min(self.__comb_forced_path[decision_idx], num_selections - 1) + if decision_idx < len(self.__comb_forced_path) + else 0 + ) + comb_chosen_selection = list(all_selections[chosen_idx]) + num_choices = len(comb_chosen_selection) + else: + num_choices = ( + self.__rng.integers(from_value, to_value, endpoint=True) if from_value < to_value else from_value + ) else: num_choices = 0 if not optional and from_value > 0: @@ -1566,11 +1695,14 @@ class TreeProcessor(lark.visitors.Interpreter): + (f" and separating with '{escape_single_quotes(separator)}'" if num_choices > 1 else ""), ) if num_choices > 0: - selected_choices: list[dict] = ( - list(self.__rng.choice(available_choices, size=num_choices, p=weights, replace=repeating)) - if available_choices - else [] - ) + if self.state.options.do_combinatorial and comb_chosen_selection is not None: + selected_choices: list[dict] = comb_chosen_selection + else: + selected_choices: list[dict] = ( + list(self.__rng.choice(available_choices, size=num_choices, p=weights, replace=repeating)) + if available_choices + else [] + ) if self.state.options.keep_choices_order: selected_choices = sorted(selected_choices, key=lambda x: x["choice_index"]) selected_choices_text = [] @@ -1829,7 +1961,7 @@ class TreeProcessor(lark.visitors.Interpreter): """ t1 = time.monotonic_ns() chosen_choices = [] - start_result = self.result + start_result = self.__result seen_wildcards_len = len(self.__seen_wildcards) applied_options = self.__clean_wildcard_options(self.__convert_choices_options(tree.children[0], False)) wildcard_key: str = self.__visit(tree.children[1], False, True) @@ -1838,8 +1970,8 @@ class TreeProcessor(lark.visitors.Interpreter): self.log(logging.DEBUG, f"Processing wildcard: {wildcard_key}") selected_wildcards = self.state.wildcards_obj.get_wildcards(wildcard_key) if not selected_wildcards: - self.detectedWildcards.append((wc, self.__is_negative)) - self.result += wc + self.__detectedWildcards.append((wc, self.__is_negative)) + self.__result += wc t2 = time.monotonic_ns() self.__debug_end("wildcard", start_result, t2 - t1, wc) return [] @@ -1891,8 +2023,8 @@ class TreeProcessor(lark.visitors.Interpreter): choice_values_all = [] for wildcard in selected_wildcards: if wildcard is None: - self.detectedWildcards.append((wc, self.__is_negative)) - self.result += wc + self.__detectedWildcards.append((wc, self.__is_negative)) + self.__result += wc t2 = time.monotonic_ns() self.__debug_end("wildcard", start_result, t2 - t1, wc) return [] @@ -1903,7 +2035,7 @@ class TreeProcessor(lark.visitors.Interpreter): continue self.__seen_wildcards.append(wildcard.key) self.log(logging.DEBUG, f"Seen wildcard '{escape_single_quotes(wildcard.key)}'") - (options, choice_values) = self.__check_wildcard_initialization(wildcard) + options, choice_values = self.__check_wildcard_initialization(wildcard) if options is not None: if applied_options is None: applied_options = options @@ -1916,7 +2048,7 @@ class TreeProcessor(lark.visitors.Interpreter): applied_options, choice_values_all, filter_specifier, wildcard_key ) if chosen_choices: - self.result += prefix + separator.join(chosen_choices) + suffix + self.__result += prefix + separator.join(chosen_choices) + suffix if wildcard_key in self.__wildcard_filters: del self.__wildcard_filters[wildcard_key] if variablename is not None: @@ -1924,8 +2056,8 @@ class TreeProcessor(lark.visitors.Interpreter): if variablebackup is not None: self.state.variables.set_user(variablename, variablebackup) elif self.state.options.if_wildcards != IFWILDCARDS_CHOICES.remove: - self.detectedWildcards.append((wc, self.__is_negative)) - self.result += wc + self.__detectedWildcards.append((wc, self.__is_negative)) + self.__result += wc if self.__debug_level == DEBUG_LEVEL.full: list_unseen = [f"'{escape_single_quotes(x)}'" for x in self.__seen_wildcards[seen_wildcards_len:]] self.log(logging.DEBUG, f"Unseen wildcards: {', '.join(list_unseen)}") @@ -1959,7 +2091,7 @@ class TreeProcessor(lark.visitors.Interpreter): Process a choices construct in the tree. """ t1 = time.monotonic_ns() - start_result = self.result + start_result = self.__result options = self.__convert_choices_options(tree.children[0], False) choice_values = [self.__convert_choice(c) for c in tree.children[1::]] ch = self.__get_original_node_content(tree, "?{...}") @@ -1967,22 +2099,22 @@ class TreeProcessor(lark.visitors.Interpreter): self.log(logging.DEBUG, "Processing choices:") prefix, chosen_choices, separator, suffix = self.__get_choices_select(options, choice_values) if chosen_choices: - self.result += prefix + separator.join(chosen_choices) + suffix + self.__result += prefix + separator.join(chosen_choices) + suffix elif self.state.options.if_wildcards != IFWILDCARDS_CHOICES.remove: - self.detectedWildcards.append((ch, self.__is_negative)) - self.result += ch + self.__detectedWildcards.append((ch, self.__is_negative)) + self.__result += ch t2 = time.monotonic_ns() self.__debug_end("choices", start_result, t2 - t1, f"'{escape_single_quotes(ch)}'") def __default__(self, tree): t1 = time.monotonic_ns() - start_result = self.result + start_result = self.__result self.__visit(tree.children) t2 = time.monotonic_ns() self.__debug_end(tree.data.value, start_result, t2 - t1) def __process_negtags(self): - # process the found negative tags + # process the found negative tags for negtag in self.__negtags: if self.state.options.cup_merge_attention: # join consecutive attention elements @@ -2032,49 +2164,49 @@ class TreeProcessor(lark.visitors.Interpreter): position = negtag.parameters or "s" if position.startswith("i"): n = int(position[1]) - self.insertion_at[n] = [negtag.start, negtag.end] + self.__insertion_at[n] = [negtag.start, negtag.end] elif len(content) > 0: if content not in self.__already_processed: if self.state.options.stn_ignore_repeats: self.__already_processed.append(content) self.log(logging.DEBUG, f"Adding content at position {position}: {content}") if position == "e": - self.add_at["end"].append(content) + self.__add_at["end"].append(content) elif position.startswith("p"): n = int(position[1]) - self.add_at["insertion_point"][n].append(content) + self.__add_at["insertion_point"][n].append(content) else: # position == "s" or invalid - self.add_at["start"].append(content) + self.__add_at["start"].append(content) else: self.log(logging.WARNING, f"Ignoring repeated content: {content}") - self.__negtags = [] + self.__negtags = [] def __apply_stn_insertions(self): """ Apply all accumulated STN content from add_at to self.result using the recorded insertion_at positions, then reset both so ppp.py does not re-apply them. """ - pos, neg = self.result.split(self.NEGATIVE_SEP, 1) + pos, neg = self.__result.split(self.NEGATIVE_SEP, 1) neg_start = len(pos) + len(self.NEGATIVE_SEP) stn_sep = self.state.options.stn_separator - self.log(logging.DEBUG, f"Applying STN additions to negative: {self.add_at}") - self.log(logging.DEBUG, f"Applying STN indexes: {self.insertion_at}") + self.log(logging.DEBUG, f"Applying STN additions to negative: {self.__add_at}") + self.log(logging.DEBUG, f"Applying STN indexes: {self.__insertion_at}") ordered_range = sorted( range(10), - key=lambda x: self.insertion_at[x][0] if self.insertion_at[x] is not None else float("-inf"), + key=lambda x: self.__insertion_at[x][0] if self.__insertion_at[x] is not None else float("-inf"), reverse=True, ) for n in ordered_range: - if self.insertion_at[n] is not None: - insertion_point_n: list[str] = self.add_at["insertion_point"][n] - ipp = self.insertion_at[n][0] - neg_start - ipl = self.insertion_at[n][1] - self.insertion_at[n][0] + if self.__insertion_at[n] is not None: + insertion_point_n: list[str] = self.__add_at["insertion_point"][n] + ipp = self.__insertion_at[n][0] - neg_start + ipl = self.__insertion_at[n][1] - self.__insertion_at[n][0] if neg[ipp - len(stn_sep) : ipp] == stn_sep: - ipp -= len(stn_sep) # adjust for existing start separator + ipp -= len(stn_sep) # adjust for existing start separator ipl += len(stn_sep) insertion_point_n.insert(0, neg[:ipp]) if neg[ipp + ipl : ipp + ipl + len(stn_sep)] == stn_sep: - ipl += len(stn_sep) # adjust for existing end separator + ipl += len(stn_sep) # adjust for existing end separator end_part = neg[ipp + ipl :] if len(end_part) > 0: insertion_point_n.append(end_part) @@ -2083,30 +2215,30 @@ class TreeProcessor(lark.visitors.Interpreter): ipp = 0 if neg.startswith(stn_sep): ipp = len(stn_sep) - self.add_at["insertion_point"][n].append(neg[ipp:]) - neg = stn_sep.join(self.add_at["insertion_point"][n]) - if self.add_at["start"]: - add_at_start = self.add_at["start"] + self.__add_at["insertion_point"][n].append(neg[ipp:]) + neg = stn_sep.join(self.__add_at["insertion_point"][n]) + if self.__add_at["start"]: + add_at_start = self.__add_at["start"] if len(neg) > 0: ipp = 0 if neg.startswith(stn_sep): ipp = len(stn_sep) # adjust for existing end separator add_at_start.append(neg[ipp:]) neg = stn_sep.join(add_at_start) - if self.add_at["end"]: - add_at_end = self.add_at["end"] + if self.__add_at["end"]: + add_at_end = self.__add_at["end"] if len(neg) > 0: ipl = len(neg) if neg.endswith(stn_sep): - ipl -= len(stn_sep) # adjust for existing start separator + ipl -= len(stn_sep) # adjust for existing start separator add_at_end.insert(0, neg[:ipl]) neg = stn_sep.join(add_at_end) # self.add_at = {"start": [], "insertion_point": [[] for _ in range(10)], "end": []} # self.insertion_at = [None for _ in range(10)] - self.result = pos + self.NEGATIVE_SEP + neg + self.__result = pos + self.NEGATIVE_SEP + neg def start(self, tree): - self.result = "" + self.__result = "" t1 = time.monotonic_ns() self.__visit(tree.children) self.__process_negtags() diff --git a/ppp_wildcards.py b/ppp_wildcards.py index 6f89af1..f1212b1 100644 --- a/ppp_wildcards.py +++ b/ppp_wildcards.py @@ -1,5 +1,6 @@ import fnmatch import os +from pathlib import Path from typing import Optional import logging import yaml @@ -82,19 +83,10 @@ class PPPWildcards: for fullpath in list(self.__wildcard_files.keys()): if fullpath != self.LOCALINPUT_FILENAME: path = os.path.dirname(fullpath) - if not os.path.exists(fullpath): + if not os.path.exists(fullpath) or not any( + Path(path).is_relative_to(folder) for folder in self.__wildcards_folders + ): self.__remove_wildcards_from_path(fullpath) - else: - a = False - for folder in self.__wildcards_folders: - try: - if os.path.commonpath([folder, path]) == folder: - a = True - break - except ValueError: - pass - if not a: - self.__remove_wildcards_from_path(fullpath) elif wildcards_input is None: self.__remove_wildcards_from_path(fullpath) if wildcards_folders is not None or wildcards_input is not None: diff --git a/scripts/ppp_script.py b/scripts/ppp_script.py index b1e6d20..68d1017 100644 --- a/scripts/ppp_script.py +++ b/scripts/ppp_script.py @@ -10,11 +10,11 @@ import numpy as np sys.path.append(str(Path(__file__).parent)) # base path for the extension -from modules import scripts, shared, script_callbacks # pylint: disable=import-error -from modules.processing import StableDiffusionProcessing # pylint: disable=import-error -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 modules import scripts, shared, script_callbacks # type: ignore +from modules.processing import StableDiffusionProcessing # type: ignore +from modules.shared import opts # type: ignore +from modules.paths import models_path # type: ignore +import gradio as gr # type: ignore from ppp import PromptPostProcessor from ppp_classes import IFWILDCARDS_CHOICES, ONWARNING_CHOICES, SUPPORTED_APPS, SUPPORTED_APPS_NAMES, PPPStateOptions from ppp_logging import DEBUG_LEVEL, PromptPostProcessorLogFactory, log @@ -140,7 +140,22 @@ class PromptPostProcessorA1111Script(scripts.Script): # show_label=True, elem_id="ppp_incremental_seed", ) - return [force_equal_seeds, unlink_seed, seed, incremental_seed] + gr.HTML("
") + with gr.Row(equal_height=True): + combinatorial = gr.Checkbox( + label="Combinatorial mode", + info="Generate all prompt combinations and cycle through them to fill the batch.", + value=PromptPostProcessor.DEFAULT_DO_COMBINATORIAL, + elem_id="ppp_combinatorial", + ) + combinatorial_limit = gr.Number( + label="Combinations limit (0 = no limit)", + value=PromptPostProcessor.DEFAULT_COMBINATORIAL_LIMIT, + precision=0, + min_width=120, + elem_id="ppp_combinatorial_limit", + ) + return [force_equal_seeds, unlink_seed, seed, incremental_seed, combinatorial, combinatorial_limit] def process( self, @@ -149,6 +164,8 @@ class PromptPostProcessorA1111Script(scripts.Script): input_unlink_seed, input_seed, input_incremental_seed, + input_combinatorial, + input_combinatorial_limit, ): # pylint: disable=arguments-differ """ Processes the prompts and applies post-processing operations. @@ -159,6 +176,8 @@ class PromptPostProcessorA1111Script(scripts.Script): input_unlink_seed (bool): Flag indicating whether to unlink the seed. input_seed (int): The seed value. input_incremental_seed (bool): Flag indicating whether to use incremental seed. + input_combinatorial (bool): Flag indicating whether to use combinatorial mode. + input_combinatorial_limit (int): Maximum number of combinations (0 = no limit). Returns: None @@ -176,6 +195,7 @@ class PromptPostProcessorA1111Script(scripts.Script): ) ) ) + num_seeds = len(getattr(p, "all_seeds", [])) options = PPPStateOptions( debug_level=DEBUG_LEVEL(getattr(opts, "ppp_gen_debug_level", PromptPostProcessor.DEFAULT_DEBUG_LEVEL)), on_warning=ONWARNING_CHOICES(getattr(opts, "ppp_gen_onwarning", PromptPostProcessor.DEFAULT_ON_WARNING)), @@ -220,6 +240,8 @@ class PromptPostProcessorA1111Script(scripts.Script): cup_remove_extranetwork_tags=getattr( opts, "ppp_rem_removeextranetworktags", PromptPostProcessor.DEFAULT_CUP_REMOVE_EXTRANETWORK_TAGS ), + do_combinatorial=input_combinatorial, + combinatorial_limit=max(num_seeds, int(input_combinatorial_limit)) if input_combinatorial else 0, ) if self.ppp_logger is None: lf = PromptPostProcessorLogFactory() @@ -251,6 +273,7 @@ class PromptPostProcessorA1111Script(scripts.Script): "PPP unlink seed": input_unlink_seed, "PPP prompt seed": input_seed, "PPP incremental seed": input_incremental_seed, + "PPP combinatorial": input_combinatorial, } ) @@ -315,7 +338,6 @@ class PromptPostProcessorA1111Script(scripts.Script): calculated_seeds: list[int] = [] if input_unlink_seed: log(self.ppp_logger, self.ppp_debug_level, logging.INFO, "Using unlinked seed") - num_seeds = len(getattr(p, "all_seeds", [])) if input_incremental_seed: first_seed = np.random.randint(0, 2**32, dtype=np.int64) if input_seed == -1 else input_seed calculated_seeds = [first_seed + i for i in range(num_seeds)] @@ -340,6 +362,7 @@ class PromptPostProcessorA1111Script(scripts.Script): # (prompt type, typeindex) -> (new positive prompt, new negative prompt) prompts_list: dict[tuple[str, int], tuple[str, str]] = {} + extra_params = {} # adds prompts regular_type = "regular" @@ -356,25 +379,49 @@ class PromptPostProcessorA1111Script(scripts.Script): if hiresfix_exists: prompts_list[(hiresfix_type, i)] = None - # processes prompts - for prompttype, typeindex in prompts_list.keys(): - log(self.ppp_logger, self.ppp_debug_level, logging.INFO, f"processing prompts ({prompttype}[{typeindex+1}])") - key = ( - (hash_fullenv, calculated_seeds[typeindex], rpr[typeindex], rnr[typeindex]) - if prompttype == regular_type - else (hash_fullenv, calculated_seeds[typeindex], rph[typeindex], rnh[typeindex]) - ) - cached = self.lru_cache.get(key) - if cached is None: - (hsh, seed, prompt, negative_prompt) = key - posp, negp, _ = ppp.process_prompt(prompt, negative_prompt, seed) - cached = (posp, negp) - self.lru_cache.put(key, cached) - # adds also the result so i2i doesn't process it unnecessarily - self.lru_cache.put((hsh, seed, posp, negp), cached) - else: - log(self.ppp_logger, self.ppp_debug_level, logging.INFO, "result already in cache") - prompts_list[(prompttype, typeindex)] = cached + if input_combinatorial: + seed_for_comb = calculated_seeds[0] if calculated_seeds else 0 + regular_copy = (rpr.copy() if rpr else None, rnr.copy() if rnr else None) + hiresfix_copy = (rph.copy() if rph else None, rnh.copy() if rnh else None) + regular_changes = False + hiresfix_changes = False + if regular_exists: + log(self.ppp_logger, self.ppp_debug_level, logging.INFO, "processing prompts combinatorially (regular)") + comb_results = ppp.process_prompt(rpr[0], rnr[0], seed_for_comb) + num_comb = len(comb_results) + for i in range(len(rpr)): # pylint: disable=consider-using-enumerate + posp, negp, _ = comb_results[i % num_comb] + prompts_list[(regular_type, i)] = (posp, negp) + extra_params["PPP combination"] = [1+(i % num_comb) for i in range(len(rpr))] + if hiresfix_exists: + log(self.ppp_logger, self.ppp_debug_level, logging.INFO, "processing prompts combinatorially (hiresfix)") + comb_results_hr = ppp.process_prompt(rph[0], rnh[0], seed_for_comb) + num_comb_hr = len(comb_results_hr) + for i in range(len(rph)): # pylint: disable=consider-using-enumerate + posp, negp, _ = comb_results_hr[i % num_comb_hr] + prompts_list[(hiresfix_type, i)] = (posp, negp) + extra_params["PPP HR combination"] = [1+(i % num_comb_hr) for i in range(len(rph))] + else: + # processes prompts + for prompttype, typeindex in prompts_list.keys(): + log(self.ppp_logger, self.ppp_debug_level, logging.INFO, f"processing prompts ({prompttype}[{typeindex+1}])") + key = ( + (hash_fullenv, calculated_seeds[typeindex], rpr[typeindex], rnr[typeindex]) + if prompttype == regular_type + else (hash_fullenv, calculated_seeds[typeindex], rph[typeindex], rnh[typeindex]) + ) + cached = self.lru_cache.get(key) + if cached is None: + (hsh, seed, prompt, negative_prompt) = key + results = ppp.process_prompt(prompt, negative_prompt, seed) + posp, negp, _ = results[0] + cached = (posp, negp) + self.lru_cache.put(key, cached) + # adds also the result so i2i doesn't process it unnecessarily + self.lru_cache.put((hsh, seed, posp, negp), cached) + else: + log(self.ppp_logger, self.ppp_debug_level, logging.INFO, "result already in cache") + prompts_list[(prompttype, typeindex)] = cached # with open(os.path.join(os.path.dirname(os.path.realpath(__file__)), "last_prompts.txt"), "w", encoding="utf-8") as f: # for (prompttype, typeindex), (posp, negp) in prompts_list.items(): @@ -404,7 +451,6 @@ class PromptPostProcessorA1111Script(scripts.Script): rnh[typeindex] = negp # initialize extra generation parameters - extra_params = {} if add_prompts: if regular_changes: extra_params["PPP original prompts"] = regular_copy[0] diff --git a/tests/base_tests.py b/tests/base_tests.py index a3bf851..817046e 100644 --- a/tests/base_tests.py +++ b/tests/base_tests.py @@ -6,10 +6,10 @@ import unittest import datetime from ppp_classes import IFWILDCARDS_CHOICES, ONWARNING_CHOICES, PPPStateOptions -from ppp_enmappings import PPPExtraNetworkMappings # pylint: disable=import-error -from ppp_wildcards import PPPWildcards # pylint: disable=import-error -from ppp import PromptPostProcessor # pylint: disable=import-error -from ppp_logging import DEBUG_LEVEL, PromptPostProcessorLogFactory # pylint: disable=import-error +from ppp_enmappings import PPPExtraNetworkMappings # type: ignore +from ppp_wildcards import PPPWildcards # type: ignore +from ppp import PromptPostProcessor # type: ignore +from ppp_logging import DEBUG_LEVEL, PromptPostProcessorLogFactory # type: ignore class PromptPair(NamedTuple): @@ -17,6 +17,12 @@ class PromptPair(NamedTuple): negative_prompt: str = "" +class OutputTuple(NamedTuple): + prompt: str = "" + negative_prompt: str = "" + variables: dict[str, str] = None + + class TestPromptPostProcessorBase(unittest.TestCase): """ A test case class for testing the PromptPostProcessor class. @@ -106,29 +112,7 @@ class TestPromptPostProcessorBase(unittest.TestCase): 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 - """ + def init_obj(self, ppp: Optional[str | PromptPostProcessor] = None) -> PromptPostProcessor: if isinstance(ppp, str): if ppp == "nocup": the_obj = PromptPostProcessor( @@ -180,30 +164,123 @@ class TestPromptPostProcessorBase(unittest.TestCase): self.wildcards_obj, self.extranetwork_maps_obj, ) + return the_obj + + def process( + self, + input_prompts: PromptPair, + expected_output: Optional[OutputTuple | list[OutputTuple]] = None, + seed: int = 1, + ppp: Optional[str | PromptPostProcessor] = None, + interrupted: bool = False, + ): + """ + Process the prompt and compare the results with the expected prompts. + + Args: + input_prompts (PromptPair): The input prompts. + expected_output (OutputTuple | list[OutputTuple], optional): The expected output. When a list is provided, the test will run once for each expected output, using the same input prompt, but seed will be incremented for each iteration. + 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. + + Returns: + None + """ + the_obj = self.init_obj(ppp) out = ( - [PromptPair("", "")] - if expected_output_prompts is None - else expected_output_prompts if isinstance(expected_output_prompts, list) else [expected_output_prompts] + [OutputTuple("", "", None)] + if expected_output is None + else expected_output if isinstance(expected_output, list) else [expected_output] ) for eo in out: - result_prompt, result_negative_prompt, output_variables = the_obj.process_prompt( + result = 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", + result_prompt, result_negative_prompt, output_variables = result[0] + if expected_output is not None: + self.assertTrue( + result_prompt == eo.prompt and result_negative_prompt == eo.negative_prompt, + f"Incorrect result '{eo.prompt}' / '{eo.negative_prompt}', got '{result_prompt}' / '{result_negative_prompt}'", + ) + if eo.variables: + unmatched_vars = {} + expected_values = {} + for var_name, var_value in eo.variables.items(): + if var_name not in output_variables or output_variables[var_name] != var_value: + unmatched_vars[var_name] = ( + output_variables[var_name] if var_name in output_variables else None + ) + expected_values[var_name] = var_value + self.assertTrue( + not unmatched_vars, + f"Result '{eo.prompt}' / '{eo.negative_prompt}' found, but variables do not match: expected {expected_values}, got {unmatched_vars}", ) seed += 1 + + def process_combinatorial( + self, + input_prompts: PromptPair, + expected_output: Optional[OutputTuple | list[OutputTuple]] = None, + seed: int = 1, + ppp: Optional[str | PromptPostProcessor] = None, + interrupted: bool = False, + ): + """ + Process the prompt and compare the results with the expected prompts. + + Args: + input_prompts (PromptPair): The input prompts. + expected_output (OutputTuple | list[OutputTuple], optional): The expected output. + 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. + + Returns: + None + """ + the_obj = self.init_obj(ppp) + out = ( + [OutputTuple("", "", None)] + if expected_output is None + else expected_output if isinstance(expected_output, list) else [expected_output] + ) + result = the_obj.process_prompt( + input_prompts.prompt, + input_prompts.negative_prompt, + seed, + combinatorial=True, + ) + self.assertEqual(self.interrupted, interrupted, "Interrupted flag is incorrect") + if not self.interrupted: + if expected_output is not None: + self.assertEqual( + len(result), len(out), f"Incorrect number of combinations (expected {len(out)}, got {len (result)})" + ) + for out_prompt, out_negative_prompt, out_variables in out: + found = None + for r_prompt, r_negative_prompt, r_variables in result: + if r_prompt == out_prompt and r_negative_prompt == out_negative_prompt: + found = OutputTuple(r_prompt, r_negative_prompt, r_variables) + break + self.assertTrue( + bool(found), + f"Combination '{out_prompt}' / '{out_negative_prompt}' not found in output", + ) + if found and out_variables: + unmatched_vars = {} + expected_values = {} + for var_name, var_value in out_variables.items(): + if var_name not in found.variables or found.variables[var_name] != var_value: + unmatched_vars[var_name] = ( + found.variables[var_name] if var_name in found.variables else None + ) + expected_values[var_name] = var_value + self.assertTrue( + not unmatched_vars, + f"Combination '{out_prompt}' / '{out_negative_prompt}' found, but variables do not match: expected {expected_values}, got {unmatched_vars}", + ) diff --git a/tests/tests_choices.py b/tests/tests_choices.py index de5f3f2..c01ae79 100644 --- a/tests/tests_choices.py +++ b/tests/tests_choices.py @@ -1,7 +1,7 @@ from dataclasses import replace -from ppp import PromptPostProcessor # pylint: disable=import-error -from .base_tests import PromptPair, TestPromptPostProcessorBase +from ppp import PromptPostProcessor # type: ignore +from .base_tests import OutputTuple, PromptPair, TestPromptPostProcessorBase if __name__ == "__main__": raise SystemExit("This script must not be run directly") @@ -17,14 +17,14 @@ class TestChoices(TestPromptPostProcessorBase): 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", ""), + OutputTuple("the choices are: choice2", ""), ppp="nocup", ) def test_ch_unsupportedsampler(self): # unsupported sampler self.process( PromptPair("the choices are: {@choice1|choice2|choice3}", ""), - PromptPair("", ""), + OutputTuple("", ""), ppp="nocup", interrupted=True, ) @@ -35,28 +35,28 @@ class TestChoices(TestPromptPostProcessorBase): "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", ""), + OutputTuple("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", ""), + OutputTuple("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", ""), + OutputTuple("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", ""), + OutputTuple("the choices are: choice1, choice3", ""), ppp="nocup", ) @@ -66,21 +66,21 @@ class TestChoices(TestPromptPostProcessorBase): "${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", ""), + OutputTuple("the choices are: choice1 choice11, choice3", ""), ppp="nocup", ) def test_ch_choicesinsidelora(self): # simple choices inside a lora self.process( PromptPair("", ""), - PromptPair("", ""), + OutputTuple("", ""), ppp="nocup", ) def test_ch_removelorawithchoices(self): self.process( PromptPair("", ""), - PromptPair("", ""), + OutputTuple("", ""), ppp=PromptPostProcessor( self.ppp_logger, self.def_env_info, @@ -98,6 +98,21 @@ class TestChoices(TestPromptPostProcessorBase): def test_ch_cmd_includewildcard(self): self.process( PromptPair("{ch_one|ch_two|%0.5::include yaml/wildcard1}", ""), - PromptPair("ch_two", ""), + OutputTuple("ch_two", ""), ppp="nocup", ) + + # Combinatorial + + def test_ch_combinatorial(self): + self.process_combinatorial( + PromptPair("{choice1|choice2|choice3}, ${v:{option1|option2}}", ""), + [ + OutputTuple("choice1, option1", ""), + OutputTuple("choice1, option2", ""), + OutputTuple("choice2, option1", ""), + OutputTuple("choice2, option2", ""), + OutputTuple("choice3, option1", ""), + OutputTuple("choice3, option2", "", {"v": "option2"}), + ], + ) diff --git a/tests/tests_cleanup.py b/tests/tests_cleanup.py index cd8e644..a05a7d0 100644 --- a/tests/tests_cleanup.py +++ b/tests/tests_cleanup.py @@ -1,8 +1,8 @@ import logging from dataclasses import replace -from ppp import PromptPostProcessor # pylint: disable=import-error -from .base_tests import PromptPair, TestPromptPostProcessorBase +from ppp import PromptPostProcessor # type: ignore +from .base_tests import OutputTuple, PromptPair, TestPromptPostProcessorBase if __name__ == "__main__": @@ -19,7 +19,7 @@ class TestCleanup(TestPromptPostProcessorBase): 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"), + OutputTuple("this is a ((test), (test,:2):1.5), (red:1.5)", "normal quality"), ) def test_cl_complex(self): # complex cleanup @@ -28,7 +28,7 @@ class TestCleanup(TestPromptPostProcessorBase): " 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( + OutputTuple( "this is BREAKABLE a (test:1.21) AND( ANDERSON (test:2):1.5) :o BREAK (red:1.5)", "[:hands, feet, :0.15]normal quality", ), @@ -37,7 +37,7 @@ class TestCleanup(TestPromptPostProcessorBase): def test_cl_removenetworktags(self): # remove network tags self.process( PromptPair("this is a test__yaml/wildcard7__", ""), - PromptPair("this is a test", ""), + OutputTuple("this is a test", ""), ppp=PromptPostProcessor( self.ppp_logger, self.def_env_info, @@ -55,7 +55,7 @@ class TestCleanup(TestPromptPostProcessorBase): 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", ""), + OutputTuple("this is a test,\nsecond line", ""), ppp=PromptPostProcessor( self.ppp_logger, self.def_env_info, @@ -79,7 +79,7 @@ class TestCleanup(TestPromptPostProcessorBase): (d:0.9)""", "", ), - PromptPair( + OutputTuple( """ (l:1.1) (d:0.9), (l:1.1) (d:0.9)""", @@ -115,7 +115,7 @@ class TestCleanup(TestPromptPostProcessorBase): "this is (a test:0.9) of (attention (merging:1.2)) where ((this)) ((is joined:1.2)) and ([this too]:1.3)", "", ), - PromptPair( + OutputTuple( "this is [a test] of (attention (merging:1.2)) where (this:1.21) (is joined:1.32) and (this too:1.17)", "", ), @@ -127,7 +127,7 @@ class TestCleanup(TestPromptPostProcessorBase): "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( + OutputTuple( "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)", "", ), @@ -140,7 +140,7 @@ class TestCleanup(TestPromptPostProcessorBase): with self.assertLogs("PromptPostProcessor", level=logging.WARNING) as cm: self.process( PromptPair("(unclosed paren", ""), - PromptPair("(unclosed paren", ""), + OutputTuple("(unclosed paren", ""), ) self.assertTrue( any("Unmatched" in msg for msg in cm.output), @@ -151,7 +151,7 @@ class TestCleanup(TestPromptPostProcessorBase): with self.assertLogs("PromptPostProcessor", level=logging.WARNING) as cm: self.process( PromptPair("extra close paren)", ""), - PromptPair("extra close paren)", ""), + OutputTuple("extra close paren)", ""), ) self.assertTrue( any("Unmatched" in msg for msg in cm.output), @@ -162,7 +162,7 @@ class TestCleanup(TestPromptPostProcessorBase): with self.assertLogs("PromptPostProcessor", level=logging.WARNING) as cm: self.process( PromptPair("(mismatched]", ""), - PromptPair("(mismatched]", ""), + OutputTuple("(mismatched]", ""), ) self.assertTrue( any("Mismatched" in msg or "Unmatched" in msg for msg in cm.output), @@ -173,7 +173,7 @@ class TestCleanup(TestPromptPostProcessorBase): with self.assertLogs("PromptPostProcessor", level=logging.WARNING) as cm: self.process( PromptPair("unclosed [bracket", ""), - PromptPair("unclosed [bracket", ""), + OutputTuple("unclosed [bracket", ""), ) self.assertTrue( any("Unmatched" in msg for msg in cm.output), @@ -184,7 +184,7 @@ class TestCleanup(TestPromptPostProcessorBase): with self.assertLogs("PromptPostProcessor", level=logging.WARNING) as cm: self.process( PromptPair("[(unmatched [bracket))", ""), - PromptPair("[(unmatched [bracket))", ""), + OutputTuple("[(unmatched [bracket))", ""), ) self.assertTrue( any("Unmatched" in msg for msg in cm.output), @@ -195,5 +195,5 @@ class TestCleanup(TestPromptPostProcessorBase): with self.assertNoLogs("PromptPostProcessor", level=logging.WARNING): self.process( PromptPair(r"text with \(escaped unmatched\]", ""), - PromptPair(r"text with \(escaped unmatched\]", ""), + OutputTuple(r"text with \(escaped unmatched\]", ""), ) diff --git a/tests/tests_host.py b/tests/tests_host.py index f38ee98..b58f6be 100644 --- a/tests/tests_host.py +++ b/tests/tests_host.py @@ -1,5 +1,5 @@ -from ppp import PromptPostProcessor # pylint: disable=import-error -from .base_tests import PromptPair, TestPromptPostProcessorBase +from ppp import PromptPostProcessor # type: ignore +from .base_tests import OutputTuple, PromptPair, TestPromptPostProcessorBase if __name__ == "__main__": raise SystemExit("This script must not be run directly") @@ -18,7 +18,7 @@ class TestHosts(TestPromptPostProcessorBase): "[test1] (test2) (test3:1.5) [(test4)]", "", ), - PromptPair("(test1:0.9) (test2) (test3:1.5) (test4:0.99)", ""), + OutputTuple("(test1:0.9) (test2) (test3:1.5) (test4:0.99)", ""), ppp=PromptPostProcessor( self.ppp_logger, { @@ -39,7 +39,7 @@ class TestHosts(TestPromptPostProcessorBase): "[test1] (test2) (test3:1.5)", "", ), - PromptPair("test1 test2 test3", ""), + OutputTuple("test1 test2 test3", ""), ppp=PromptPostProcessor( self.ppp_logger, { @@ -60,7 +60,7 @@ class TestHosts(TestPromptPostProcessorBase): "[test1] (test2) (test3:1.5)", "", ), - PromptPair("", ""), + OutputTuple("", ""), ppp=PromptPostProcessor( self.ppp_logger, { @@ -81,7 +81,7 @@ class TestHosts(TestPromptPostProcessorBase): "[test1] (test2) (test3:1.5)", "", ), - PromptPair("", ""), + OutputTuple("", ""), ppp=PromptPostProcessor( self.ppp_logger, { @@ -103,7 +103,7 @@ class TestHosts(TestPromptPostProcessorBase): "[test1:test2:0.5]", "", ), - PromptPair("test1", ""), + OutputTuple("test1", ""), ppp=PromptPostProcessor( self.ppp_logger, { @@ -124,7 +124,7 @@ class TestHosts(TestPromptPostProcessorBase): "[test1:test2:0.5]", "", ), - PromptPair("test2", ""), + OutputTuple("test2", ""), ppp=PromptPostProcessor( self.ppp_logger, { @@ -145,7 +145,7 @@ class TestHosts(TestPromptPostProcessorBase): "[test1::0.5] [:test2:0.5] [test3:test4:0.5]", "", ), - PromptPair("test1 test3", ""), + OutputTuple("test1 test3", ""), ppp=PromptPostProcessor( self.ppp_logger, { @@ -166,7 +166,7 @@ class TestHosts(TestPromptPostProcessorBase): "[test1:test2:0.5]", "", ), - PromptPair("", ""), + OutputTuple("", ""), ppp=PromptPostProcessor( self.ppp_logger, { @@ -187,7 +187,7 @@ class TestHosts(TestPromptPostProcessorBase): "[test1:test2:0.5]", "", ), - PromptPair("", ""), + OutputTuple("", ""), ppp=PromptPostProcessor( self.ppp_logger, { @@ -209,7 +209,7 @@ class TestHosts(TestPromptPostProcessorBase): "[test1|test2|test3]", "", ), - PromptPair("test1", ""), + OutputTuple("test1", ""), ppp=PromptPostProcessor( self.ppp_logger, { @@ -230,7 +230,7 @@ class TestHosts(TestPromptPostProcessorBase): "[test1|test2|test3]", "", ), - PromptPair("", ""), + OutputTuple("", ""), ppp=PromptPostProcessor( self.ppp_logger, { @@ -251,7 +251,7 @@ class TestHosts(TestPromptPostProcessorBase): "[test1|test2|test3]", "", ), - PromptPair("", ""), + OutputTuple("", ""), ppp=PromptPostProcessor( self.ppp_logger, { @@ -273,7 +273,7 @@ class TestHosts(TestPromptPostProcessorBase): "test1 AND test2:2", "", ), - PromptPair("test1\ntest2", ""), + OutputTuple("test1\ntest2", ""), ppp=PromptPostProcessor( self.ppp_logger, { @@ -294,7 +294,7 @@ class TestHosts(TestPromptPostProcessorBase): "test1 AND test2:2", "", ), - PromptPair("test1, test2", ""), + OutputTuple("test1, test2", ""), ppp=PromptPostProcessor( self.ppp_logger, { @@ -315,7 +315,7 @@ class TestHosts(TestPromptPostProcessorBase): "test1 AND test2:2", "", ), - PromptPair("test1 test2", ""), + OutputTuple("test1 test2", ""), ppp=PromptPostProcessor( self.ppp_logger, { @@ -336,7 +336,7 @@ class TestHosts(TestPromptPostProcessorBase): "test1 AND test2:2", "", ), - PromptPair("", ""), + OutputTuple("", ""), ppp=PromptPostProcessor( self.ppp_logger, { @@ -358,7 +358,7 @@ class TestHosts(TestPromptPostProcessorBase): "test1 BREAK test2", "", ), - PromptPair("test1\ntest2", ""), + OutputTuple("test1\ntest2", ""), ppp=PromptPostProcessor( self.ppp_logger, { @@ -379,7 +379,7 @@ class TestHosts(TestPromptPostProcessorBase): "test1 BREAK test2", "", ), - PromptPair("test1, test2", ""), + OutputTuple("test1, test2", ""), ppp=PromptPostProcessor( self.ppp_logger, { @@ -400,7 +400,7 @@ class TestHosts(TestPromptPostProcessorBase): "test1 BREAK test2", "", ), - PromptPair("test1 test2", ""), + OutputTuple("test1 test2", ""), ppp=PromptPostProcessor( self.ppp_logger, { @@ -421,7 +421,7 @@ class TestHosts(TestPromptPostProcessorBase): "test1 BREAK test2", "", ), - PromptPair("", ""), + OutputTuple("", ""), ppp=PromptPostProcessor( self.ppp_logger, { diff --git a/tests/tests_stn.py b/tests/tests_stn.py index c755ffa..5cc2e40 100644 --- a/tests/tests_stn.py +++ b/tests/tests_stn.py @@ -1,4 +1,4 @@ -from .base_tests import PromptPair, TestPromptPostProcessorBase +from .base_tests import OutputTuple, PromptPair, TestPromptPostProcessorBase if __name__ == "__main__": raise SystemExit("This script must not be run directly") @@ -17,7 +17,7 @@ class TestSendToNegative(TestPromptPostProcessorBase): "flowersred, green, blueyellow, purpleblack", "normal quality, worse quality", ), - PromptPair("flowers", "red, green, yellow, normal quality, purple, worse quality, black, blue"), + OutputTuple("flowers", "red, green, yellow, normal quality, purple, worse quality, black, blue"), ) def test_stn_complex(self): # complex negtags @@ -26,7 +26,7 @@ class TestSendToNegative(TestPromptPostProcessorBase): "red ((pink)), flowers purple, mauveblue, yellow green", "normal quality, , bad quality, worse quality", ), - PromptPair( + OutputTuple( "flowers", "red, (pink:1.21), normal quality, mauve, yellow, bad quality, green, worse quality, purple, blue", ), @@ -38,7 +38,7 @@ class TestSendToNegative(TestPromptPostProcessorBase): "red ((pink)), flowers purple, mauveblue, yellow green", "normal quality, , bad quality, worse quality", ), - PromptPair( + OutputTuple( " (()), flowers , , ", "red, ((pink)), normal quality, mauve, yellow, bad quality, green, worse quality, purple, blue", ), @@ -51,7 +51,8 @@ class TestSendToNegative(TestPromptPostProcessorBase): "[neg1] this is a ((testneg2) (test:2.0): 1.5 ) (red[square]:1.5)", "normal quality", ), - PromptPair( + OutputTuple( + "this is a ((test) (test:2):1.5) (red:1.5)", "[neg1], ([square]:1.5), normal quality, (neg2:1.65)" ), ) @@ -62,7 +63,7 @@ class TestSendToNegative(TestPromptPostProcessorBase): "this is a (([complexneg1|simpleneg2|regularneg3] test)(test:2.0):1.5)", "normal quality", ), - PromptPair( + OutputTuple( "this is a (([complex|simple|regular] test)(test:2):1.5)", "([neg1||]:1.65), ([|neg2|]:1.65), ([||neg3]:1.65), normal quality", ), @@ -74,7 +75,7 @@ class TestSendToNegative(TestPromptPostProcessorBase): "this is a (([complexneg1[one|twoneg12||three|four(neg14)]|simpleneg2|regularneg3] test)(test:2.0):1.5)", "normal quality", ), - PromptPair( + OutputTuple( "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", ), @@ -83,7 +84,7 @@ class TestSendToNegative(TestPromptPostProcessorBase): 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]")], + OutputTuple("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 @@ -92,7 +93,7 @@ class TestSendToNegative(TestPromptPostProcessorBase): "[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( + OutputTuple( "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]", ), @@ -104,7 +105,7 @@ class TestSendToNegative(TestPromptPostProcessorBase): "[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( + OutputTuple( "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]", ), @@ -116,7 +117,7 @@ class TestSendToNegative(TestPromptPostProcessorBase): "[pos1neg1[pos11|pos12neg12||pos14|pos15neg15]|pos2neg2|pos3neg3]", "", ), - PromptPair( + OutputTuple( "[pos1[pos11|pos12||pos14|pos15]|pos2|pos3]", "[neg1||], [[|neg12|||]||], [[||||neg15]||], [|neg2|], [||neg3]", # "[neg1[|neg12|||neg15]|neg2|neg3]", # expected output if the constructs were unified diff --git a/tests/tests_varcomms.py b/tests/tests_varcomms.py index 535cb01..d85bb3e 100644 --- a/tests/tests_varcomms.py +++ b/tests/tests_varcomms.py @@ -1,8 +1,8 @@ from dataclasses import replace -from ppp import PromptPostProcessor # pylint: disable=import-error -from ppp_classes import ONWARNING_CHOICES # pylint: disable=import-error -from .base_tests import PromptPair, TestPromptPostProcessorBase +from ppp import PromptPostProcessor # type: ignore +from ppp_classes import ONWARNING_CHOICES # type: ignore +from .base_tests import OutputTuple, PromptPair, TestPromptPostProcessorBase if __name__ == "__main__": raise SystemExit("This script must not be run directly") @@ -21,8 +21,7 @@ class TestVarCommands(TestPromptPostProcessorBase): "${v1=}${v3:}", "", ), - PromptPair("", ""), - variables={"v1": "", "v2": "", "v3": ""}, + OutputTuple("", "",{"v1": "", "v2": "", "v3": ""}), ) # Echoed variables tests @@ -34,8 +33,7 @@ class TestVarCommands(TestPromptPostProcessorBase): "", ), # v3 is echoed withs two defaults, the output prompt has both but the variable value is the last default - PromptPair("test3test4", ""), - variables={"v1": "test1", "v2": "test2", "v3": "test4"}, + OutputTuple("test3test4", "", {"v1": "test1", "v2": "test2", "v3": "test4"}), ) def test_unknown_echoed_variable(self): @@ -44,8 +42,7 @@ class TestVarCommands(TestPromptPostProcessorBase): "${v1}", "", ), - PromptPair("", ""), - variables={"v1": ""}, + OutputTuple("", "", {"v1": ""}), ppp=PromptPostProcessor( self.ppp_logger, self.def_env_info, @@ -68,7 +65,7 @@ class TestVarCommands(TestPromptPostProcessorBase): "${f=filename}${w=0.5}", "", ), - PromptPair("", ""), + OutputTuple("", ""), ) # Variable nesting tests @@ -79,8 +76,7 @@ class TestVarCommands(TestPromptPostProcessorBase): "${v1=test ${v2:OK}}${v1}", "", ), - PromptPair("test OK", ""), - variables={"v1": "test OK", "v2": "OK"}, + OutputTuple("test OK", "", {"v1": "test OK", "v2": "OK"}), ) def test_var_nested_2(self): # variable set nested in variable default @@ -89,8 +85,7 @@ class TestVarCommands(TestPromptPostProcessorBase): "${v1:test ${v2=OK}${v2}}", "", ), - PromptPair("test OK", ""), - variables={"v1": "test OK", "v2": "OK"}, + OutputTuple("test OK", "", {"v1": "test OK", "v2": "OK"}), ) def test_var_nested_3(self): # variable default nested in variable default @@ -99,8 +94,7 @@ class TestVarCommands(TestPromptPostProcessorBase): "${v1:test ${v2:OK}}", "", ), - PromptPair("test OK", ""), - variables={"v1": "test OK", "v2": "OK"}, + OutputTuple("test OK", "", {"v1": "test OK", "v2": "OK"}), ) # Array variable tests @@ -111,8 +105,7 @@ class TestVarCommands(TestPromptPostProcessorBase): "${v1[]=val1}${v1[]+=val2}${v1[]+=val3}${v1[1]:defval},${v1[]:defval2},${v1[&'.']:defval3}", "", ), - PromptPair("val2,val1, val2, val3,val1.val2.val3", ""), - variables={"v1[]": "val1, val2, val3", "v1[1]": "val2", "v1[&'.']": "val1.val2.val3"}, + OutputTuple("val2,val1, val2, val3,val1.val2.val3", "", {"v1[]": "val1, val2, val3", "v1[1]": "val2", "v1[&'.']": "val1.val2.val3"}), ) def test_array_variable_2(self): # override of array variable value, test of default value when array variable is empty, test of default value when array variable is not set @@ -121,8 +114,7 @@ class TestVarCommands(TestPromptPostProcessorBase): "${v1[]=val1}${v1[]=val2}${v1[]:defval},${v2[]:defval2},${v2[1]:defval3},${v3[]=}${v3[]:defval4}", "", ), - PromptPair("val2,defval2,defval3", ""), - variables={"v1[]": "val2", "v2[]": "defval2", "v2[1]": "defval3", "v3[]": ""}, + OutputTuple("val2,defval2,defval3", "", {"v1[]": "val2", "v2[]": "defval2", "v2[1]": "defval3", "v3[]": ""}), ) def test_array_variable_3(self): # access array index by variable, set array variable to expanded array variable and add expanded array @@ -131,8 +123,7 @@ class TestVarCommands(TestPromptPostProcessorBase): "${v1[]=val1}${v1[]+=val2}${v2=1}${v1[v2]:defval1}${v3[]=${v1[]}}${v3[]+=${v1[]}}, ${v3[&'.']}", "", ), - PromptPair("val2, val1, val2.val1, val2", ""), - variables={"v1[]": "val1, val2", "v2": "1", "v1[v2]": "val2", "v3[]": "val1, val2, val1, val2", "v3[&'.']": "val1, val2.val1, val2"}, + OutputTuple("val2, val1, val2.val1, val2", "", {"v1[]": "val1, val2", "v2": "1", "v1[v2]": "val2", "v3[]": "val1, val2, val1, val2", "v3[&'.']": "val1, val2.val1, val2"}), ) def test_array_variable_4(self): # test list in array @@ -141,8 +132,7 @@ class TestVarCommands(TestPromptPostProcessorBase): "${v1[]=val1}${v1[]+=val2}${v1[]+=val3}OKnot OK", "", ), - PromptPair("OK", ""), - variables={"v1[]": "val1, val2, val3"}, + OutputTuple("OK", "", {"v1[]": "val1, val2, val3"}), ) def test_array_variable_5(self): # test empty array @@ -151,8 +141,7 @@ class TestVarCommands(TestPromptPostProcessorBase): "${v1[]=}OKnot OK,OKnot OK", "", ), - PromptPair("OK,OK", ""), - variables={"v1[]": ""}, + OutputTuple("OK,OK", "", {"v1[]": ""}), ) def test_array_variable_6(self): # array variable set and addition with expanded values from array variables @@ -161,8 +150,7 @@ class TestVarCommands(TestPromptPostProcessorBase): "${v1[]=val1}${v1[]+=val2}${v2[]=val3}${v3[]=*v1[]}${v3[]+=*v2[]}", "", ), - PromptPair("", ""), - variables={"v1[]": "val1, val2", "v2[]": "val3", "v3[]": "val1, val2, val3"}, + OutputTuple("", "", {"v1[]": "val1, val2", "v2[]": "val3", "v3[]": "val1, val2, val3"}), ) def test_array_variable_7(self): # array variable set and addition with expanded values from wildcards @@ -171,8 +159,7 @@ class TestVarCommands(TestPromptPostProcessorBase): "${v1[]=*__yaml/wildcard1__}${v1[]+=*__yaml/wildcard2__}${v1[2]:defval}", "", ), - PromptPair("choice3", ""), - variables={"v1[]": "choice2, choice1, choice3, choice1"}, + OutputTuple("choice3", "", {"v1[]": "choice2, choice1, choice3, choice1"}), ) def test_array_variable_8(self): # array variable set and addition with expanded values from lists @@ -181,8 +168,7 @@ class TestVarCommands(TestPromptPostProcessorBase): "${v1[]=*()}${v1[]+=*('one','two')}${v2=three}${v1[]+=*(v2,'four')}${v1[2]:defval}", "", ), - PromptPair("three", ""), - variables={"v1[]": "one, two, three, four", "v2": "three"}, + OutputTuple("three", "", {"v1[]": "one, two, three, four", "v2": "three"}), ) def test_array_variable_9(self): # array variable length @@ -191,8 +177,7 @@ class TestVarCommands(TestPromptPostProcessorBase): "${v1[]=val1}${v1[]+=val2}${v1[]+=val3}${v1[#]:defval}, OKnot OK", "", ), - PromptPair("3, OK", ""), - variables={"v1[]": "val1, val2, val3", "v1[#]": "3"}, + OutputTuple("3, OK", "", {"v1[]": "val1, val2, val3", "v1[#]": "3"}), ) def test_array_variable_10(self): # array variable set with expanded values from wildcards in command format @@ -201,8 +186,7 @@ class TestVarCommands(TestPromptPostProcessorBase): "*__yaml/wildcard1__*__yaml/wildcard2__${v1[2]:defval}", "", ), - PromptPair("choice3", ""), - variables={"v1[]": "choice2, choice1, choice3, choice1"}, + OutputTuple("choice3", "", {"v1[]": "choice2, choice1, choice3, choice1"}), ) # Operator tests @@ -215,7 +199,7 @@ class TestVarCommands(TestPromptPostProcessorBase): "${r1=hello}${r2=hello}OKnot OK", "", ), - PromptPair("OK", ""), + OutputTuple("OK", ""), ) def test_operator_RnoteqR(self): # test for the not before the operator @@ -224,7 +208,7 @@ class TestVarCommands(TestPromptPostProcessorBase): "${r1=hello}${r2=bye}OKnot OK", "", ), - PromptPair("OK", ""), + OutputTuple("OK", ""), ) def test_operator_RneR(self): @@ -233,7 +217,7 @@ class TestVarCommands(TestPromptPostProcessorBase): "${r1=hello}${r2=bye}OKnot OK", "", ), - PromptPair("OK", ""), + OutputTuple("OK", ""), ) def test_operator_RltR(self): @@ -242,7 +226,7 @@ class TestVarCommands(TestPromptPostProcessorBase): "${r1=1}${r2=2}OKnot OK", "", ), - PromptPair("OK", ""), + OutputTuple("OK", ""), ) def test_operator_RgtR(self): @@ -251,7 +235,7 @@ class TestVarCommands(TestPromptPostProcessorBase): "${r1=2}${r2=1}OKnot OK", "", ), - PromptPair("OK", ""), + OutputTuple("OK", ""), ) def test_operator_RleR(self): @@ -260,7 +244,7 @@ class TestVarCommands(TestPromptPostProcessorBase): "${r1=1}${r2=1}OKnot OK", "", ), - PromptPair("OK", ""), + OutputTuple("OK", ""), ) def test_operator_RgeR(self): @@ -269,7 +253,7 @@ class TestVarCommands(TestPromptPostProcessorBase): "${r1=1}${r2=1}OKnot OK", "", ), - PromptPair("OK", ""), + OutputTuple("OK", ""), ) def test_operator_RinR(self): @@ -278,7 +262,7 @@ class TestVarCommands(TestPromptPostProcessorBase): "${r1=hello}${r2=hello world}OKnot OK", "", ), - PromptPair("OK", ""), + OutputTuple("OK", ""), ) def test_operator_RcontainsR(self): @@ -287,7 +271,7 @@ class TestVarCommands(TestPromptPostProcessorBase): "${r1=hello world}${r2=hello}OKnot OK", "", ), - PromptPair("OK", ""), + OutputTuple("OK", ""), ) ## A vs A @@ -298,7 +282,7 @@ class TestVarCommands(TestPromptPostProcessorBase): "${a1[]=*('hello','world')}${a2[]=*('hello','world')}OKnot OK", "", ), - PromptPair("OK", ""), + OutputTuple("OK", ""), ) def test_operator_AneA_1(self): @@ -307,7 +291,7 @@ class TestVarCommands(TestPromptPostProcessorBase): "${a1[]=*('hello')}${a2[]=*('bye')}OKnot OK", "", ), - PromptPair("OK", ""), + OutputTuple("OK", ""), ) def test_operator_AneA_2(self): @@ -316,7 +300,7 @@ class TestVarCommands(TestPromptPostProcessorBase): "${a1[]=*('hello','world')}${a2[]=*('hello')}OKnot OK", "", ), - PromptPair("OK", ""), + OutputTuple("OK", ""), ) def test_operator_AneA_3(self): @@ -325,7 +309,7 @@ class TestVarCommands(TestPromptPostProcessorBase): "${a1[]=*('hello','world')}${a2[]=*('world','hello')}OKnot OK", "", ), - PromptPair("OK", ""), + OutputTuple("OK", ""), ) def test_operator_AltA_1(self): @@ -334,7 +318,7 @@ class TestVarCommands(TestPromptPostProcessorBase): "${a1[]=*(1,2,3)}${a2[]=*(2,3,4)}OKnot OK", "", ), - PromptPair("OK", ""), + OutputTuple("OK", ""), ) def test_operator_AltA_2(self): @@ -343,7 +327,7 @@ class TestVarCommands(TestPromptPostProcessorBase): "${a1[]=*(1,2)}${a2[]=*(2,3,4)}OKnot OK", "", ), - PromptPair("not OK", ""), + OutputTuple("not OK", ""), ) def test_operator_AgtA(self): @@ -352,7 +336,7 @@ class TestVarCommands(TestPromptPostProcessorBase): "${a1[]=*(2,3,4)}${a2[]=*(1,2,3)}OKnot OK", "", ), - PromptPair("OK", ""), + OutputTuple("OK", ""), ) def test_operator_AleA(self): @@ -361,7 +345,7 @@ class TestVarCommands(TestPromptPostProcessorBase): "${a1[]=*(1,2)}${a2[]=*(1,3)}OKnot OK", "", ), - PromptPair("OK", ""), + OutputTuple("OK", ""), ) def test_operator_AgeA(self): @@ -370,7 +354,7 @@ class TestVarCommands(TestPromptPostProcessorBase): "${a1[]=*(1,3)}${a2[]=*(1,2)}OKnot OK", "", ), - PromptPair("OK", ""), + OutputTuple("OK", ""), ) def test_operator_AinA(self): @@ -379,7 +363,7 @@ class TestVarCommands(TestPromptPostProcessorBase): "${a1[]=*('hello')}${a2[]=*('hello', 'world')}OKnot OK", "", ), - PromptPair("OK", ""), + OutputTuple("OK", ""), ) def test_operator_AcontainsA(self): @@ -388,7 +372,7 @@ class TestVarCommands(TestPromptPostProcessorBase): "${a1[]=*('hello','world')}${a2[]=*('hello')}OKnot OK", "", ), - PromptPair("OK", ""), + OutputTuple("OK", ""), ) ## A vs R @@ -399,7 +383,7 @@ class TestVarCommands(TestPromptPostProcessorBase): "${a1[]=*('hello','world')}${r2=hello}OKnot OK", "", ), - PromptPair("not OK", ""), + OutputTuple("not OK", ""), ppp="nostrict", ) @@ -409,7 +393,7 @@ class TestVarCommands(TestPromptPostProcessorBase): "${a1[]=*('hello')}${r2=bye}OKnot OK", "", ), - PromptPair("OK", ""), + OutputTuple("OK", ""), ppp="nostrict", ) @@ -419,7 +403,7 @@ class TestVarCommands(TestPromptPostProcessorBase): "${a1[]=*(1,2,3)}${r2=2}OKnot OK", "", ), - PromptPair("not OK", ""), + OutputTuple("not OK", ""), ppp="nostrict", ) @@ -429,7 +413,7 @@ class TestVarCommands(TestPromptPostProcessorBase): "${a1[]=*(2,3,4)}${r2=2}OKnot OK", "", ), - PromptPair("not OK", ""), + OutputTuple("not OK", ""), ppp="nostrict", ) @@ -439,7 +423,7 @@ class TestVarCommands(TestPromptPostProcessorBase): "${a1[]=*(1,2)}${r2=2}OKnot OK", "", ), - PromptPair("OK", ""), + OutputTuple("OK", ""), ppp="nostrict", ) @@ -449,7 +433,7 @@ class TestVarCommands(TestPromptPostProcessorBase): "${a1[]=*(1,3)}${r2=2}OKnot OK", "", ), - PromptPair("not OK", ""), + OutputTuple("not OK", ""), ppp="nostrict", ) @@ -459,7 +443,7 @@ class TestVarCommands(TestPromptPostProcessorBase): "${a1[]=*('hello', 'world')}${r2=hello world)}OKnot OK", "", ), - PromptPair("OK", ""), + OutputTuple("OK", ""), ) def test_operator_AcontainsR(self): @@ -468,7 +452,7 @@ class TestVarCommands(TestPromptPostProcessorBase): "${a1[]=*('hello','world')}${r2=hello}OKnot OK", "", ), - PromptPair("OK", ""), + OutputTuple("OK", ""), ) ## R vs A @@ -479,7 +463,7 @@ class TestVarCommands(TestPromptPostProcessorBase): "${r1=hello}${a2[]=*('hello','world')}OKnot OK", "", ), - PromptPair("not OK", ""), + OutputTuple("not OK", ""), ppp="nostrict", ) @@ -489,7 +473,7 @@ class TestVarCommands(TestPromptPostProcessorBase): "${r1=bye}${a2[]=*('hello')}OKnot OK", "", ), - PromptPair("OK", ""), + OutputTuple("OK", ""), ppp="nostrict", ) @@ -499,7 +483,7 @@ class TestVarCommands(TestPromptPostProcessorBase): "${r1=2}${a2[]=*(1,2,3)}OKnot OK", "", ), - PromptPair("not OK", ""), + OutputTuple("not OK", ""), ppp="nostrict", ) @@ -509,7 +493,7 @@ class TestVarCommands(TestPromptPostProcessorBase): "${r1=2}${a2[]=*(2,3,4)}OKnot OK", "", ), - PromptPair("not OK", ""), + OutputTuple("not OK", ""), ppp="nostrict", ) @@ -519,7 +503,7 @@ class TestVarCommands(TestPromptPostProcessorBase): "${r1=2}${a2[]=*(1,2)}OKnot OK", "", ), - PromptPair("not OK", ""), + OutputTuple("not OK", ""), ppp="nostrict", ) @@ -529,7 +513,7 @@ class TestVarCommands(TestPromptPostProcessorBase): "${r1=2}${a2[]=*(1,3)}OKnot OK", "", ), - PromptPair("not OK", ""), + OutputTuple("not OK", ""), ppp="nostrict", ) @@ -539,7 +523,7 @@ class TestVarCommands(TestPromptPostProcessorBase): "${r1=hello}${a2[]=*('hello', 'world')}OKnot OK", "", ), - PromptPair("OK", ""), + OutputTuple("OK", ""), ) def test_operator_RcontainsA(self): @@ -548,7 +532,7 @@ class TestVarCommands(TestPromptPostProcessorBase): "${r1=hello world}${a2[]=*('hello','world')}OKnot OK", "", ), - PromptPair("OK", ""), + OutputTuple("OK", ""), ) ## R vs V @@ -559,7 +543,7 @@ class TestVarCommands(TestPromptPostProcessorBase): "${r1=hello}OKnot OK", "", ), - PromptPair("OK", ""), + OutputTuple("OK", ""), ) def test_operator_ReqV_str_fail(self): @@ -568,7 +552,7 @@ class TestVarCommands(TestPromptPostProcessorBase): "${r1=hello}OKnot OK", "", ), - PromptPair("not OK", ""), + OutputTuple("not OK", ""), interrupted=True, ) @@ -578,7 +562,7 @@ class TestVarCommands(TestPromptPostProcessorBase): "${r1=42}OKnot OK", "", ), - PromptPair("OK", ""), + OutputTuple("OK", ""), ) def test_operator_ReqV_num_fail(self): @@ -587,7 +571,7 @@ class TestVarCommands(TestPromptPostProcessorBase): "${r1=42}OKnot OK", "", ), - PromptPair("not OK", ""), + OutputTuple("not OK", ""), interrupted=True, ) @@ -597,7 +581,7 @@ class TestVarCommands(TestPromptPostProcessorBase): "${r1=true}OKnot OK", "", ), - PromptPair("OK", ""), + OutputTuple("OK", ""), ) def test_operator_ReqV_bool_fail(self): @@ -606,7 +590,7 @@ class TestVarCommands(TestPromptPostProcessorBase): "${r1=true}OKnot OK", "", ), - PromptPair("not OK", ""), + OutputTuple("not OK", ""), interrupted=True, ) @@ -618,7 +602,7 @@ class TestVarCommands(TestPromptPostProcessorBase): "${a1[]=*('hello','world')}OKnot OK", "", ), - PromptPair("OK", ""), + OutputTuple("OK", ""), ) def test_listoperand_LinA(self): @@ -627,7 +611,7 @@ class TestVarCommands(TestPromptPostProcessorBase): "${a2[]=*('hello','world')}OKnot OK", "", ), - PromptPair("OK", ""), + OutputTuple("OK", ""), ) # Indexed operands @@ -638,7 +622,7 @@ class TestVarCommands(TestPromptPostProcessorBase): "${a2[]=*('hello','world')}OKnot OK", "", ), - PromptPair("OK", ""), + OutputTuple("OK", ""), ) # Float values @@ -649,7 +633,7 @@ class TestVarCommands(TestPromptPostProcessorBase): "${a=1.5}OKnot OK", "", ), - PromptPair("OK", ""), + OutputTuple("OK", ""), ) @@ -661,7 +645,7 @@ class TestVarCommands(TestPromptPostProcessorBase): "OKnot OK", "", ), - PromptPair("not OK", ""), + OutputTuple("not OK", ""), ppp=PromptPostProcessor( self.ppp_logger, self.def_env_info, @@ -682,7 +666,7 @@ class TestVarCommands(TestPromptPostProcessorBase): "OKnot OK", "", ), - PromptPair("", ""), + OutputTuple("", ""), interrupted=True, ) @@ -692,7 +676,7 @@ class TestVarCommands(TestPromptPostProcessorBase): "abcOKnot OK", "", ), - PromptPair("not OK", ""), + OutputTuple("not OK", ""), ppp=PromptPostProcessor( self.ppp_logger, self.def_env_info, @@ -713,7 +697,7 @@ class TestVarCommands(TestPromptPostProcessorBase): "abcOKnot OK", "", ), - PromptPair("", ""), + OutputTuple("", ""), interrupted=True, ) @@ -723,7 +707,7 @@ class TestVarCommands(TestPromptPostProcessorBase): "OKnot OK", "", ), - PromptPair("not OK", ""), + OutputTuple("not OK", ""), ppp=PromptPostProcessor( self.ppp_logger, self.def_env_info, @@ -746,7 +730,7 @@ class TestVarCommands(TestPromptPostProcessorBase): "[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( + OutputTuple( "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]", ), @@ -758,7 +742,7 @@ class TestVarCommands(TestPromptPostProcessorBase): "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( + OutputTuple( "this \\(is\\): a (([complex|simple|regular] test)(test:2):1.5)\nBREAK :0.5 AND hypernettrigger :0.3", "normal quality", ), @@ -770,7 +754,7 @@ class TestVarCommands(TestPromptPostProcessorBase): "this is SD1PONYSD2NOPONYNOPONY", "", ), - PromptPair("this is PONY", ""), + OutputTuple("this is PONY", ""), ppp=PromptPostProcessor( self.ppp_logger, { @@ -788,19 +772,19 @@ class TestVarCommands(TestPromptPostProcessorBase): def test_cmd_set_if(self): # set and if commands self.process( PromptPair("valuethis test is OKnot OK", ""), - PromptPair("this test is OK", ""), + OutputTuple("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", ""), + OutputTuple("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", ""), + OutputTuple("this test is OK", ""), ) def test_cmd_set_if_echo_nested(self): # nested set, if and echo commands @@ -809,7 +793,7 @@ class TestVarCommands(TestPromptPostProcessorBase): "1OKnot OK NOK OK", "", ), - PromptPair("OK OK OK", ""), + OutputTuple("OK OK OK", ""), ) def test_cmd_set_if_complex_conditions_1(self): # complex conditions (or) @@ -818,7 +802,7 @@ class TestVarCommands(TestPromptPostProcessorBase): "truefalsethis test is OKnot OK", "", ), - PromptPair("this test is OK", ""), + OutputTuple("this test is OK", ""), ) def test_cmd_set_if_complex_conditions_2(self): # complex conditions (and) @@ -827,13 +811,13 @@ class TestVarCommands(TestPromptPostProcessorBase): "truetruethis test is OKnot OK", "", ), - PromptPair("this test is OK", ""), + OutputTuple("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", ""), + OutputTuple("this test is OK", ""), ) def test_cmd_set_if_complex_conditions_4(self): # complex conditions (not, precedence) @@ -842,7 +826,7 @@ class TestVarCommands(TestPromptPostProcessorBase): "truefalsethis test is OKnot OK", "", ), - PromptPair("this test is OK", ""), + OutputTuple("this test is OK", ""), ) def test_cmd_set_if_complex_conditions_5(self): # complex conditions (not, precedence, comparison) @@ -851,7 +835,7 @@ class TestVarCommands(TestPromptPostProcessorBase): "1falsethis test is OKnot OK", "", ), - PromptPair("this test is OK", ""), + OutputTuple("this test is OK", ""), ) def test_cmd_set_if_complex_conditions_6(self): # complex conditions @@ -860,7 +844,7 @@ class TestVarCommands(TestPromptPostProcessorBase): "123this test is OKnot OK", "", ), - PromptPair("this test is OK", ""), + OutputTuple("this test is OK", ""), ) def test_cmd_set_if_complex_conditions_7(self): # complex conditions @@ -869,7 +853,7 @@ class TestVarCommands(TestPromptPostProcessorBase): "123this test is OKnot OK", "", ), - PromptPair("this test is OK", ""), + OutputTuple("this test is OK", ""), ) def test_cmd_set_if2(self): # set and more complex if commands @@ -878,7 +862,7 @@ class TestVarCommands(TestPromptPostProcessorBase): "First: value1this test is OKOK2not OK\nSecond: value3this test is OKnot OK", "", ), - PromptPair("First: this test is OK\nSecond: this test is OK", ""), + OutputTuple("First: this test is OK\nSecond: this test is OK", ""), ) def test_cmd_set_add_if(self): # set, add and if commands @@ -887,7 +871,7 @@ class TestVarCommands(TestPromptPostProcessorBase): "value2this test is OKnot OK", "", ), - PromptPair("this test is OK", ""), + OutputTuple("this test is OK", ""), ) def test_cmd_set_add_DP_if(self): # set, add (DP format) and if commands @@ -896,7 +880,7 @@ class TestVarCommands(TestPromptPostProcessorBase): "${v=value}${v+=2}this test is OKnot OK", "", ), - PromptPair("this test is OK", ""), + OutputTuple("this test is OK", ""), ) def test_cmd_set_immediateeval(self): # set (DP format) with mixed evaluation @@ -905,7 +889,7 @@ class TestVarCommands(TestPromptPostProcessorBase): "${var=!__yaml/wildcard1__}the choices are: ${var}, ${var}, ${var2:default}, ${var3=__yaml/wildcard1__}${var3}, ${var3}", "", ), - PromptPair("the choices are: choice2, choice2, default, choice3, choice1", ""), + OutputTuple("the choices are: choice2, choice2, default, choice3, choice1", ""), ppp="nocup", ) @@ -915,7 +899,7 @@ class TestVarCommands(TestPromptPostProcessorBase): "${var=__yaml/wildcard1__}the choices are: ${var}, ${var}, ${var+=, __yaml/wildcard2__}${var}, ${var}, ${var+=!, __yaml/wildcard3__}${var}, ${var}", "", ), - PromptPair( + OutputTuple( "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 ", "", ), @@ -928,7 +912,7 @@ class TestVarCommands(TestPromptPostProcessorBase): "valuethis test is OKnot OK", "", ), - PromptPair("this test is OK", ""), + OutputTuple("this test is OK", ""), ) def test_cmd_set_ifundefined_if_2(self): # set, ifundefined and if commands @@ -937,7 +921,7 @@ class TestVarCommands(TestPromptPostProcessorBase): "valuevalue2this test is OKnot OK", "", ), - PromptPair("this test is OK", ""), + OutputTuple("this test is OK", ""), ) def test_cmd_set_ifundefined_DP_if(self): # set, ifundefined (DP format) and if commands @@ -946,7 +930,7 @@ class TestVarCommands(TestPromptPostProcessorBase): "${v?=value}this test is OKnot OK", "", ), - PromptPair("this test is OK", ""), + OutputTuple("this test is OK", ""), ) def test_cmd_set_ifundefined_DP_if_2(self): # set, ifundefined (DP format) and if commands @@ -955,7 +939,7 @@ class TestVarCommands(TestPromptPostProcessorBase): "${v=!value}${v?=!value2}this test is OKnot OK", "", ), - PromptPair("this test is OK", ""), + OutputTuple("this test is OK", ""), ) def test_cmd_echo_sysvar(self): @@ -964,7 +948,7 @@ class TestVarCommands(TestPromptPostProcessorBase): "${_model:defval}", "", ), - PromptPair("sdxl", ""), + OutputTuple("sdxl", ""), ) def test_cmd_ext(self): # ext @@ -973,7 +957,7 @@ class TestVarCommands(TestPromptPostProcessorBase): "trigger1trigger2trigger4trigger5", "", ), - PromptPair( + OutputTuple( "trigger1,trigger2,trigger4,trigger5", "", ), @@ -985,7 +969,7 @@ class TestVarCommands(TestPromptPostProcessorBase): "", "", ), - PromptPair("triggergeneric1, triggergeneric2, two, triggergeneric1, triggergeneric2, two", ""), + OutputTuple("triggergeneric1, triggergeneric2, two, triggergeneric1, triggergeneric2, two", ""), ) def test_cmd_ext_map1(self): # ext mapping, no lora @@ -994,7 +978,7 @@ class TestVarCommands(TestPromptPostProcessorBase): "inlinetrigger", "", ), - PromptPair("inlinetrigger, triggergeneric1, triggergeneric2, two", ""), + OutputTuple("inlinetrigger, triggergeneric1, triggergeneric2, two", ""), ) def test_cmd_ext_map2(self): # ext mapping, lora with weight @@ -1003,7 +987,7 @@ class TestVarCommands(TestPromptPostProcessorBase): "inlinetrigger", "", ), - PromptPair("inlinetrigger, triggerpony1, triggerpony2", ""), + OutputTuple("inlinetrigger, triggerpony1, triggerpony2", ""), ppp=PromptPostProcessor( self.ppp_logger, { @@ -1024,7 +1008,7 @@ class TestVarCommands(TestPromptPostProcessorBase): "inlinetrigger", "", ), - PromptPair("inlinetrigger, triggerpony1, triggerpony2", ""), + OutputTuple("inlinetrigger, triggerpony1, triggerpony2", ""), ppp=PromptPostProcessor( self.ppp_logger, { @@ -1045,7 +1029,7 @@ class TestVarCommands(TestPromptPostProcessorBase): "inlinetrigger", "", ), - PromptPair("inlinetrigger, triggerpony1, triggerpony2", ""), + OutputTuple("inlinetrigger, triggerpony1, triggerpony2", ""), ppp=PromptPostProcessor( self.ppp_logger, { @@ -1066,7 +1050,7 @@ class TestVarCommands(TestPromptPostProcessorBase): "inlinetrigger", "", ), - PromptPair("inlinetrigger, triggerillustrious1, triggerillustrious2", ""), + OutputTuple("inlinetrigger, triggerillustrious1, triggerillustrious2", ""), ppp=PromptPostProcessor( self.ppp_logger, { diff --git a/tests/tests_variants.py b/tests/tests_variants.py index 444c839..ac06990 100644 --- a/tests/tests_variants.py +++ b/tests/tests_variants.py @@ -1,8 +1,8 @@ from dataclasses import replace from ppp import PromptPostProcessor -from ppp_classes import ONWARNING_CHOICES # pylint: disable=import-error -from .base_tests import PromptPair, TestPromptPostProcessorBase +from ppp_classes import ONWARNING_CHOICES # type: ignore +from .base_tests import OutputTuple, PromptPair, TestPromptPostProcessorBase if __name__ == "__main__": raise SystemExit("This script must not be run directly") @@ -21,7 +21,7 @@ class TestModelVariants(TestPromptPostProcessorBase): "test1test2test3test4", "", ), - PromptPair("test1test2", ""), + OutputTuple("test1test2", ""), ppp=PromptPostProcessor( self.ppp_logger, { diff --git a/tests/tests_wildcards.py b/tests/tests_wildcards.py index a304334..ca03c0a 100644 --- a/tests/tests_wildcards.py +++ b/tests/tests_wildcards.py @@ -2,7 +2,7 @@ from dataclasses import replace from ppp import PromptPostProcessor from ppp_classes import IFWILDCARDS_CHOICES -from .base_tests import PromptPair, TestPromptPostProcessorBase +from .base_tests import OutputTuple, PromptPair, TestPromptPostProcessorBase if __name__ == "__main__": raise SystemExit("This script must not be run directly") @@ -18,7 +18,7 @@ class TestWildcards(TestPromptPostProcessorBase): def test_wc_ignore(self): # wildcards with ignore option self.process( PromptPair("__bad_wildcard__", "{option1|option2}"), - PromptPair("__bad_wildcard__", "{option1|option2}"), + OutputTuple("__bad_wildcard__", "{option1|option2}"), ppp=PromptPostProcessor( self.ppp_logger, self.def_env_info, @@ -40,7 +40,7 @@ class TestWildcards(TestPromptPostProcessorBase): "[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( + OutputTuple( "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]", ), @@ -62,7 +62,7 @@ class TestWildcards(TestPromptPostProcessorBase): def test_wc_warn(self): # wildcards with warn option self.process( PromptPair("__bad_wildcard__", "{option1|option2}"), - PromptPair(PromptPostProcessor.WILDCARD_WARNING + "__bad_wildcard__", "{option1|option2}"), + OutputTuple(PromptPostProcessor.WILDCARD_WARNING + "__bad_wildcard__", "{option1|option2}"), ppp=PromptPostProcessor( self.ppp_logger, self.def_env_info, @@ -81,7 +81,7 @@ class TestWildcards(TestPromptPostProcessorBase): def test_wc_stop(self): # wildcards with stop option self.process( PromptPair("__bad_wildcard__", "{option1|option2}"), - PromptPair( + OutputTuple( PromptPostProcessor.WILDCARD_STOP.format("__bad_wildcard__") + "__bad_wildcard__", "{option1|option2}", ), @@ -104,7 +104,7 @@ class TestWildcards(TestPromptPostProcessorBase): def test_wcinvar_warn(self): # wildcards in var with warn option self.process( PromptPair("${v=__bad_wildcard__}${v}", ""), - PromptPair(PromptPostProcessor.WILDCARD_WARNING + "__bad_wildcard__", ""), + OutputTuple(PromptPostProcessor.WILDCARD_WARNING + "__bad_wildcard__", ""), ppp=PromptPostProcessor( self.ppp_logger, self.def_env_info, @@ -123,154 +123,154 @@ class TestWildcards(TestPromptPostProcessorBase): def test_wc_invalid_name(self): self.process( PromptPair("the choices are: ___invalid__", ""), - PromptPair("the choices are: ___invalid__", ""), + OutputTuple("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", ""), + OutputTuple("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", ""), + OutputTuple("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", ""), + OutputTuple("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", ""), + OutputTuple("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", ""), + OutputTuple("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", ""), + OutputTuple("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", ""), + OutputTuple("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", ""), + OutputTuple("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", ""), + OutputTuple("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", ""), + OutputTuple("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", ""), + OutputTuple("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", ""), + OutputTuple("the choice is: choice3-choice3", ""), ppp="nocup", ) def test_wc_wildcard_filter_index_range(self): # wildcard with positional index range filter self.process( PromptPair("the choice is: __yaml/wildcard2'2-3'__", ""), - PromptPair("the choice is: choice3-choice3", ""), + OutputTuple("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", ""), + OutputTuple("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", ""), + OutputTuple("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", ""), + OutputTuple("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", ""), + OutputTuple("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-choice1", ""), + OutputTuple("the choice is: choice3-choice1", ""), 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", ""), + OutputTuple("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: choice1bisbis", ""), + OutputTuple("the choice is: choice1bisbis", ""), ppp="nocup", ) def test_wc_wildcard_filter_compound4(self): # wildcard with doubly inherited compound filter with variable self.process( PromptPair("${v=label1}the choice is: __yaml/wildcard2bisbis'#${v}+label3'__", ""), - PromptPair("the choice is: choice1bisbis", ""), + OutputTuple("the choice is: choice1bisbis", ""), ppp="nocup", ) @@ -280,7 +280,7 @@ class TestWildcards(TestPromptPostProcessorBase): "the choice is: __yaml/wildcard2__, __yaml/wildcard2__", "", ), - PromptPair("the choice is: choice3-choice1, choice3-choice1- choice2 ", ""), + OutputTuple("the choice is: choice3-choice1, choice3-choice1- choice2 ", ""), ppp="nocup", ) @@ -290,49 +290,49 @@ class TestWildcards(TestPromptPostProcessorBase): "${v=label1}the choice is: __yaml/wildcard2__, __yaml/wildcard2__", "", ), - PromptPair("the choice is: choice3-choice1, choice3-choice1- choice2 ", ""), + OutputTuple("the choice is: choice3-choice1, 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", ""), + OutputTuple("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", ""), + OutputTuple("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", ""), + OutputTuple("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: ", ""), + OutputTuple("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", ""), + OutputTuple("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", ""), + OutputTuple("the choices are: choice2", ""), ppp="nocup", ) @@ -340,9 +340,9 @@ class TestWildcards(TestPromptPostProcessorBase): 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", ""), + OutputTuple("the choices are: choice0", ""), + OutputTuple("the choices are: choice1, choice3", ""), + OutputTuple("the choices are: choice1, choice3", ""), ], ppp="nocup", ) @@ -350,7 +350,7 @@ class TestWildcards(TestPromptPostProcessorBase): def test_wc_unsupportedsampler(self): # unsupported sampler self.process( PromptPair("the choices are: __@yaml/wildcard2__", ""), - PromptPair("", ""), + OutputTuple("", ""), ppp="nocup", interrupted=True, ) @@ -358,42 +358,42 @@ class TestWildcards(TestPromptPostProcessorBase): 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", ""), + OutputTuple("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", ""), + OutputTuple("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", ""), + OutputTuple("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", ""), + OutputTuple("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", ""), + OutputTuple("the choices are: choice2", ""), ppp="nocup", ) def test_wc_circular(self): # wildcard circular reference self.process( PromptPair("the choices are: __yaml/circular1__", ""), - PromptPair("", ""), + OutputTuple("", ""), ppp="nocup", interrupted=True, ) @@ -401,14 +401,14 @@ class TestWildcards(TestPromptPostProcessorBase): def test_wc_including(self): # wildcard including another wildcard self.process( PromptPair("the choices are: __yaml/including__", ""), - PromptPair("the choices are: choice4", ""), + OutputTuple("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("", ""), + OutputTuple("", ""), ppp="nocup", interrupted=True, ) @@ -419,6 +419,136 @@ class TestWildcards(TestPromptPostProcessorBase): "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", ""), + OutputTuple("the choices are: choice1-choice3-choice1 choice3- choice2 - choice2 choice3", ""), + ppp="nocup", + ) + + # Combinatorial + + def test_wc_combinatorial_1(self): # combinatorial wildcard with variable + self.process_combinatorial( + PromptPair("the choices are: __2$$yaml/wildcard2__, ${v:{option1|option2}}", ""), + [ # 12 combinations + OutputTuple("the choices are: choice1, choice2, option1", "", {"v": "option1"}), + OutputTuple("the choices are: choice1, choice2, option2", "", {"v": "option2"}), + OutputTuple("the choices are: choice1, choice3, option1", "", {"v": "option1"}), + OutputTuple("the choices are: choice1, choice3, option2", "", {"v": "option2"}), + OutputTuple("the choices are: choice2, choice1, option1", "", {"v": "option1"}), + OutputTuple("the choices are: choice2, choice1, option2", "", {"v": "option2"}), + OutputTuple("the choices are: choice2, choice3, option1", "", {"v": "option1"}), + OutputTuple("the choices are: choice2, choice3, option2", "", {"v": "option2"}), + OutputTuple("the choices are: choice3, choice1, option1", "", {"v": "option1"}), + OutputTuple("the choices are: choice3, choice1, option2", "", {"v": "option2"}), + OutputTuple("the choices are: choice3, choice2, option1", "", {"v": "option1"}), + OutputTuple("the choices are: choice3, choice2, option2", "", {"v": "option2"}), + ], + ) + + def test_wc_combinatorial_2(self): # combinatorial wildcard + self.process_combinatorial( + PromptPair("__yaml/wildcard2__", ""), + [ # 36 combinations + # groups of 3 + ## same choice repeated 3 times + OutputTuple("choice1-choice1-choice1", ""), + OutputTuple(" choice2 - choice2 - choice2 ", ""), + OutputTuple("choice3-choice3-choice3", ""), + ## one choice repeated 2 times in all positions + OutputTuple("choice1-choice1- choice2 ", ""), + OutputTuple("choice1-choice1-choice3", ""), + OutputTuple(" choice2 - choice2 -choice1", ""), + OutputTuple(" choice2 - choice2 -choice3", ""), + OutputTuple("choice3-choice3-choice1", ""), + OutputTuple("choice3-choice3- choice2 ", ""), + OutputTuple(" choice2 -choice1-choice1", ""), + OutputTuple("choice3-choice1-choice1", ""), + OutputTuple("choice1- choice2 - choice2 ", ""), + OutputTuple("choice3- choice2 - choice2 ", ""), + OutputTuple("choice1-choice3-choice3", ""), + OutputTuple(" choice2 -choice3-choice3", ""), + OutputTuple("choice1- choice2 -choice1", ""), + OutputTuple("choice1-choice3-choice1", ""), + OutputTuple(" choice2 -choice1- choice2 ", ""), + OutputTuple(" choice2 -choice3- choice2 ", ""), + OutputTuple("choice3-choice1-choice3", ""), + OutputTuple("choice3- choice2 -choice3", ""), + ## choices 1, 2, 3 in all positions + OutputTuple("choice1- choice2 -choice3", ""), + OutputTuple("choice1-choice3- choice2 ", ""), + OutputTuple(" choice2 -choice1-choice3", ""), + OutputTuple(" choice2 -choice3-choice1", ""), + OutputTuple("choice3-choice1- choice2 ", ""), + OutputTuple("choice3- choice2 -choice1", ""), + # groups of 2 + ## same choice repeated 2 times + OutputTuple("choice1-choice1", ""), + OutputTuple(" choice2 - choice2 ", ""), + OutputTuple("choice3-choice3", ""), + ## choices 1 and 2 in all positions + OutputTuple("choice1- choice2 ", ""), + OutputTuple(" choice2 -choice1", ""), + ## choices 2 and 3 in all positions + OutputTuple(" choice2 -choice3", ""), + OutputTuple("choice3- choice2 ", ""), + ## choices 1 and 3 in all positions + OutputTuple("choice1-choice3", ""), + OutputTuple("choice3-choice1", ""), + ], + ppp="nocup", + ) + + def test_wc_combinatorial_3(self): # combinatorial wildcard (keep choice order) + self.process_combinatorial( + PromptPair("__2-3$$-$$yaml/wildcard2__", ""), + [ # 4 combinations + # groups of 3 + ## choices 1, 2, 3 + OutputTuple("choice1- choice2 -choice3", ""), + # groups of 2 + ## choices 1 and 2 + OutputTuple("choice1- choice2 ", ""), + ## choices 2 and 3 + OutputTuple(" choice2 -choice3", ""), + ## choices 1 and 3 + OutputTuple("choice1-choice3", ""), + ], + ppp=PromptPostProcessor( + self.ppp_logger, + self.def_env_info, + replace( + self.defopts, + keep_choices_order=True, + cup_do_cleanup=False, + ), + self.grammar_content, + self.interrupt, + self.wildcards_obj, + self.extranetwork_maps_obj, + ), + ) + + def test_wc_combinatorial_4(self): # combinatorial wildcard (don't keep choice order) + self.process_combinatorial( + PromptPair("__2-3$$-$$yaml/wildcard2__", ""), + [ # 12 combinations + # groups of 3 + ## choices 1, 2, 3 in all positions + OutputTuple("choice1- choice2 -choice3", ""), + OutputTuple("choice1-choice3- choice2 ", ""), + OutputTuple(" choice2 -choice1-choice3", ""), + OutputTuple(" choice2 -choice3-choice1", ""), + OutputTuple("choice3-choice1- choice2 ", ""), + OutputTuple("choice3- choice2 -choice1", ""), + # groups of 2 + ## choices 1 and 2 in all positions + OutputTuple("choice1- choice2 ", ""), + OutputTuple(" choice2 -choice1", ""), + ## choices 2 and 3 in all positions + OutputTuple(" choice2 -choice3", ""), + OutputTuple("choice3- choice2 ", ""), + ## choices 1 and 3 in all positions + OutputTuple("choice1-choice3", ""), + OutputTuple("choice3-choice1", ""), + ], ppp="nocup", )