From 04119b169e06252d43583cb12a567ae2f6674fd7 Mon Sep 17 00:00:00 2001 From: Antonio Cordero Balcazar Date: Sat, 2 May 2026 11:18:27 +0200 Subject: [PATCH] * Improved combinatorial limit warning. * Added combinatorial shuffle. * A1111: Improved combinatorial with hiresfix. * Some refactoring. * Update of AI test instructions. Co-authored-by: Copilot --- .github/instructions/tests.instructions.md | 66 +++++-- .gitignore | 3 +- README.md | 2 +- ppp.py | 13 +- ppp_cache.py | 2 +- ppp_classes.py | 17 +- ppp_comfyui.py | 30 +++- ppp_tree.py | 19 +- ppp_variables.py | 6 +- scripts/ppp_script.py | 70 +++++--- tests/base_tests.py | 16 +- tests/tests_choices.py | 24 +-- tests/tests_cleanup.py | 28 +-- tests/tests_host.py | 42 ++--- tests/tests_performance.py | 16 +- tests/tests_stn.py | 22 +-- tests/tests_varcomms.py | 200 ++++++++++----------- tests/tests_variants.py | 4 +- tests/tests_wildcards.py | 104 +++++------ 19 files changed, 393 insertions(+), 291 deletions(-) diff --git a/.github/instructions/tests.instructions.md b/.github/instructions/tests.instructions.md index 89a011b..ac48c69 100644 --- a/.github/instructions/tests.instructions.md +++ b/.github/instructions/tests.instructions.md @@ -31,33 +31,59 @@ Examples: `test_cl_simple`, `test_ch_choices`, `test_wc_ignore` ## Running a Test Case via `self.process()` -Use the `process()` helper from the base class — never instantiate `PromptPostProcessor` directly in test methods. +Use the `process()` helper from the base class — never instantiate `PromptPostProcessor` directly in test methods unless there is a need to pass specific options not covered by the default setup or the "nocup" or "nostrict" options. ```python def test_cl_simple(self): """simple cleanup""" self.process( - "input prompt", # positive prompt - "", # negative prompt - PromptPair("expected output", ""), # expected result + InputTuple( + "input prompt", # positive prompt + ""), # negative prompt + OutputTuple( + "expected output", # expected positive prompt + "", # expected negative prompt + {} # expected variables (optional) + ), ) + +def test_cl_combinatorial(self): + """simple cleanup""" + self.process( + InputTuple( + "input prompt", # positive prompt + ""), # negative prompt + [ + OutputTuple( + "expected output", # expected positive prompt + "", # expected negative prompt + {} # expected variables (optional) + ), + OutputTuple( + "expected output", # expected positive prompt + "", # expected negative prompt + {} # expected variables (optional) + ), + ], + combinatorial=True, + ) + ``` ### `process()` Signature (key parameters) | Parameter | Type | Notes | |-----------|------|-------| -| `input_prompt` | `str` | Positive prompt input | -| `input_negative_prompt` | `str` | Negative prompt input | -| `expected_output` | `PromptPair \| list[PromptPair]` | Single or multiple valid outputs | +| `input_prompts` | `InputTuple` | Prompts input | +| `expected_output` | `OutputTuple \| list[OutputTuple]` | Single or multiple valid outputs | | `seed` | `int` | Optional, defaults to fixed seed | -| `ppp` | `PromptPostProcessor \| str \| None` | Pass `"nocup"` to skip creation | +| `ppp` | `PromptPostProcessor \| str \| None` | Supported values `"nocup"`, `"nostrict"` or a specific instance | | `interrupted` | `bool` | Expected interrupt flag | -| `output_variables` | `dict[str, str] \| None` | Variables to validate after processing | +| `combinatorial` | `bool` | Whether to run a combinatorial generation. If a specific ppp instance is used then it is ignored | ## Assertions -Use `assertEqual` with a descriptive message string: +Use `assertEqual` or similar methods with a descriptive message string: ```python self.assertEqual(result, expected, "Descriptive failure message") @@ -73,8 +99,24 @@ Override `self.defopts` or `self.def_env_info` to pass non-default options — d ```python def test_cl_custom(self): """cleanup with custom separator""" - opts = {**self.defopts, "ppp_stn_separator": " | "} - self.process("a, , b", "", PromptPair("a | b", ""), ppp_opts=opts) + self.process( + InputTuple("a, , b", ""), + OutputTuple("a | b", ""), + ppp=PromptPostProcessor( + self.ppp_logger, + self.def_env_info, + replace( + self.defopts, + keep_choices_order=True, + cup_do_cleanup=False, + do_combinatorial=True, + ), + self.grammar_content, + self.interrupt, + self.wildcards_obj, + self.extranetwork_maps_obj, + ), + ) ``` ## Entry Point diff --git a/.gitignore b/.gitignore index c17d6eb..0651593 100644 --- a/.gitignore +++ b/.gitignore @@ -6,8 +6,7 @@ venv !.vscode/settings.json !.vscode/launch.json +logs tests/tests_local.py tests/local_wildcards tests/logs - -scripts/last_prompts.txt diff --git a/README.md b/README.md index f76bd7f..3c90f10 100644 --- a/README.md +++ b/README.md @@ -78,7 +78,7 @@ 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. +To develop, I suggest doing so with the extension isolated from the UI (you can use a symlink to test it in the UI), and with its own virtual environment (venv or .venv), so the tests work and can be debugged properly. ## License diff --git a/ppp.py b/ppp.py index e632d6e..f825e3f 100644 --- a/ppp.py +++ b/ppp.py @@ -84,6 +84,7 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in DEFAULT_CUP_REMOVE_EXTRANETWORK_TAGS = defopt["cup_remove_extranetwork_tags"] DEFAULT_STRICT_OPERATORS = defopt["strict_operators"] DEFAULT_DO_COMBINATORIAL = defopt["do_combinatorial"] + DEFAULT_COMBINATORIAL_SHUFFLE = defopt["combinatorial_shuffle"] DEFAULT_COMBINATORIAL_LIMIT = defopt["combinatorial_limit"] WILDCARD_WARNING = '(WARNING TEXT "INVALID WILDCARD" IN BRIGHT RED:1.5)\nBREAK ' WILDCARD_STOP = "INVALID WILDCARD! {0}\nBREAK " @@ -882,7 +883,9 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in v.update(variables) return prompt, negative_prompt, v - def __processprompts(self, rng, prompt, negative_prompt) -> list[tuple[str, str, dict[str, str | None]]]: + def __processprompts( + self, rng: np.random.Generator, prompt: str, negative_prompt: str + ) -> list[tuple[str, str, dict[str, str | None]]]: """ Process the prompt and negative prompt. @@ -914,6 +917,7 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in try: results = processor.start_visit(parsed) except PPPInterrupt as e: + results = [] self.log(logging.ERROR, e.message) if e.pos_prefix: prompt = e.pos_prefix + prompt @@ -931,6 +935,9 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in final_results.append(self.__postprocess_result(r)) if self.state.options.do_combinatorial: self.log(logging.INFO, f"Total combinations: {len(final_results)}") + if self.state.options.combinatorial_shuffle: + rng.shuffle(final_results) + self.log(logging.INFO, "Combinations shuffled") return final_results def process_prompt( @@ -974,7 +981,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, {}] + 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, {}] + return [(original_prompt, original_negative_prompt, {})] diff --git a/ppp_cache.py b/ppp_cache.py index 44ec45f..2e15cdf 100644 --- a/ppp_cache.py +++ b/ppp_cache.py @@ -7,7 +7,7 @@ from ppp_logging import DEBUG_LEVEL class PPPLRUCache: - ProcessInput = Tuple[int, int, str, str] # (seed, wildcards_hash, positive_prompt, negative_prompt) + ProcessInput = Tuple[int, int, str, str] # (env_hash, seed, positive_prompt, negative_prompt) ProcessResult = Tuple[str, str] # (positive_prompt, negative_prompt) def __init__(self, capacity: int, logger: Logger = None, debug_level: DEBUG_LEVEL = DEBUG_LEVEL.none): diff --git a/ppp_classes.py b/ppp_classes.py index d8905e0..b4ed3db 100644 --- a/ppp_classes.py +++ b/ppp_classes.py @@ -79,7 +79,7 @@ class ModelDetectConfig(BaseModel): @model_validator(mode="after") def check_class_or_property(self) -> "ModelDetectConfig": if self.class_ is None and self.property is None: - raise ValueError("either 'class' or 'property' must be specified") + raise ValueError("Either 'class' or 'property' must be specified") return self @@ -101,17 +101,17 @@ class FindInFilenamePattern(BaseModel): flag_value = 0 for flag in v: if not isinstance(flag, str) or not hasattr(re, flag): - raise ValueError(f"invalid regex flag '{flag}'") + raise ValueError(f"Invalid regex flag '{flag}'") flag_value |= getattr(re, flag) return flag_value - raise ValueError(f"expected int or list of flag-name strings, got {type(v).__name__}") + raise ValueError(f"Expected int or list of flag-name strings, got {type(v).__name__}") @model_validator(mode="after") def validate_regex(self) -> "FindInFilenamePattern": try: re.compile(self.regex, self.flags) except re.error as exc: - raise ValueError(f"invalid regex pattern '{self.regex}': {exc}") from exc + raise ValueError(f"Invalid regex pattern '{self.regex}': {exc}") from exc return self @@ -136,9 +136,9 @@ class VariantConfig(BaseModel): elif isinstance(item, dict): normalized.append(item) else: - raise ValueError(f"expected str or dict in 'find_in_filename' list, got {type(item).__name__}") + raise ValueError(f"Expected str or dict in 'find_in_filename' list, got {type(item).__name__}") return normalized - raise ValueError(f"expected str, dict, or list for 'find_in_filename', got {type(v).__name__}") + raise ValueError(f"Expected str, dict, or list for 'find_in_filename', got {type(v).__name__}") # ------------------- Model configuration ------------------- @@ -153,7 +153,7 @@ class ModelConfig(BaseModel): @model_validator(mode="after") def check_detect_or_variants(self) -> "ModelConfig": if self.detect is None and self.variants is None: - raise ValueError("at least one of 'detect' or 'variants' must be specified") + raise ValueError("At least one of 'detect' or 'variants' must be specified") return self @@ -169,7 +169,7 @@ class PPPConfig(BaseModel): @model_validator(mode="after") def check_hosts_or_models(self) -> "PPPConfig": if self.hosts is None and self.models is None: - raise ValueError("at least one of 'hosts' or 'models' must be specified") + raise ValueError("At least one of 'hosts' or 'models' must be specified") return self # ------------------- State object ------------------- @@ -202,6 +202,7 @@ class PPPStateOptions: cup_remove_extranetwork_tags: bool = False strict_operators: bool = True do_combinatorial: bool = False + combinatorial_shuffle: bool = False combinatorial_limit: int = 100 # 0 = no limit def __post_init__(self): diff --git a/ppp_comfyui.py b/ppp_comfyui.py index 1fb7571..d3aa4d9 100644 --- a/ppp_comfyui.py +++ b/ppp_comfyui.py @@ -189,6 +189,15 @@ class PromptPostProcessorComfyUINode: "label_off": "No", }, ), + "combinatorial_shuffle": ( + "BOOLEAN", + { + "default": PromptPostProcessor.DEFAULT_COMBINATORIAL_SHUFFLE, + "tooltip": "Shuffle the combinatorial results", + "label_on": "Yes", + "label_off": "No", + }, + ), "combinatorial_limit": ( "INT", { @@ -277,6 +286,7 @@ class PromptPostProcessorComfyUINode: do_cleanup, cleanup_variables, do_combinatorial, + combinatorial_shuffle, combinatorial_limit, wc_options=None, stn_options=None, @@ -371,6 +381,7 @@ class PromptPostProcessorComfyUINode: else PromptPostProcessor.DEFAULT_CUP_REMOVE_EXTRANETWORK_TAGS ), do_combinatorial=do_combinatorial, + combinatorial_shuffle=combinatorial_shuffle, combinatorial_limit=combinatorial_limit, ) self.wildcards_obj.refresh_wildcards( @@ -393,12 +404,19 @@ class PromptPostProcessorComfyUINode: self.extranetwork_mappings_obj, ) 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, - # ) + + # with open(os.path.join(os.path.dirname(os.path.realpath(__file__)), "logs", "last_prompts_comfyui.txt"), "w", encoding="utf-8") as f: + # f.write(f"Seed: {seed if seed is not None else 1}\n") + # f.write(f"In Positive: {pos_prompt}\n") + # f.write(f"In Negative: {neg_prompt}\n") + # f.write("\n") + # for i, (posp, negp, var) in enumerate(results): + # f.write(f"Index: {i}\n") + # f.write(f"Out Positive: {posp}\n") + # f.write(f"Out Negative: {negp}\n") + # f.write(f"Out Variables: {var}\n") + # f.write("\n") + return tuple(zip(*results)) # unzip the list of tuples into tuple of lists def interrupt(self): diff --git a/ppp_tree.py b/ppp_tree.py index aa731db..0768333 100644 --- a/ppp_tree.py +++ b/ppp_tree.py @@ -124,8 +124,12 @@ class TreeProcessor(lark.visitors.Interpreter): ) return tuple(self.__comb_trace) + limit_reached = False + def _dfs(forced_path: tuple[int, ...]): + nonlocal limit_reached if 0 < limit <= len(results): + limit_reached = True return trace = _run(forced_path) # For each decision that was reached but not forced, spawn branches for all @@ -133,19 +137,21 @@ class TreeProcessor(lark.visitors.Interpreter): # 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): + limit_reached = True return num_options = trace[i] for opt in range(1, num_options): if 0 < limit <= len(results): + limit_reached = True 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): + if limit_reached: self.log( - logging.WARNING, f"Combinatorial limit of {limit} reached; some combinations may have been skipped." + logging.WARNING, f"Combinatorial limit of {limit} reached; some combinations have been skipped." ) return results @@ -1354,25 +1360,26 @@ class TreeProcessor(lark.visitors.Interpreter): enmapping = self.state.extranetwork_mappings_obj.extranetwork_mappings.get(extnet_id, None) if enmapping: for v in enmapping.variants: - if str(v.condition): + cond = str(v.condition) if v.condition is not None else None + if cond: try: cnd = parse_prompt( self.state, "condition", - str(v.condition), + cond, self.state.parsers["condition"], True, ) except lark.exceptions.UnexpectedInput as e: self.warn_or_stop( - f"Error parsing condition '{escape_single_quotes(str(v.condition))}' in extranetwork mapping '{escape_single_quotes(extnet_id)}'! : {e.__class__.__name__}", + f"Error parsing condition '{escape_single_quotes(cond)}' in extranetwork mapping '{escape_single_quotes(extnet_id)}'! : {e.__class__.__name__}", e, ) cnd = None else: cnd = "True" if cnd is not None and (cnd == "True" or self.__eval_condition(cnd)): - if str(v.condition): + if cond: found_mappings.append(v) else: else_mapping = v diff --git a/ppp_variables.py b/ppp_variables.py index 5c3e784..b5c23fc 100644 --- a/ppp_variables.py +++ b/ppp_variables.py @@ -32,14 +32,14 @@ class VariableRepository: def set_system(self, name: str, value: Any) -> None: """Set a system variable.""" if not self.name_is_system(name): - raise ValueError(f"invalid system variable name '{name}': must start with an underscore") + raise ValueError(f"Invalid system variable name '{name}': must start with an underscore") self._system[name] = value def update_system(self, mapping: dict[str, Any]) -> None: """Bulk-update system variables from *mapping*.""" for name in mapping: if not self.name_is_system(name): - raise ValueError(f"invalid system variable name '{name}': must start with an underscore") + raise ValueError(f"Invalid system variable name '{name}': must start with an underscore") self._system.update(mapping) def clear_system(self) -> None: @@ -59,7 +59,7 @@ class VariableRepository: def set_user(self, name: str, value: Any) -> None: """Set a user variable.""" if self.name_is_system(name): - raise ValueError(f"invalid user variable name '{name}': must not start with an underscore") + raise ValueError(f"Invalid user variable name '{name}': must not start with an underscore") self._user[name] = value def delete_user(self, name: str) -> None: diff --git a/scripts/ppp_script.py b/scripts/ppp_script.py index 68d1017..ae0dffe 100644 --- a/scripts/ppp_script.py +++ b/scripts/ppp_script.py @@ -103,8 +103,7 @@ class PromptPostProcessorA1111Script(scripts.Script): elem_id="ppp_force_equal_seeds", ) gr.HTML("
") - gr.Markdown( - """ + gr.Markdown(""" Unlink the seed to use the specified one for the prompts instead of the image seed. * A seed of -1 and "Incremental seed" checked will use a random seed for the first prompt and consecutive values for the rest. This is the same as when you use -1 for the image seed. @@ -113,8 +112,7 @@ class PromptPostProcessorA1111Script(scripts.Script): * Any other seed value and "Incremental seed" unchecked will use the specified seed for all the prompts. Seeds are only used for the wildcards and choice constructs. - """ - ) + """) gr.HTML("
") with gr.Row(equal_height=True): unlink_seed = gr.Checkbox( @@ -148,6 +146,12 @@ class PromptPostProcessorA1111Script(scripts.Script): value=PromptPostProcessor.DEFAULT_DO_COMBINATORIAL, elem_id="ppp_combinatorial", ) + combinatorial_shuffle = gr.Checkbox( + label="Shuffle combinations", + info="Shuffle the combinatorial results.", + value=PromptPostProcessor.DEFAULT_COMBINATORIAL_SHUFFLE, + elem_id="ppp_combinatorial_shuffle", + ) combinatorial_limit = gr.Number( label="Combinations limit (0 = no limit)", value=PromptPostProcessor.DEFAULT_COMBINATORIAL_LIMIT, @@ -155,7 +159,7 @@ class PromptPostProcessorA1111Script(scripts.Script): min_width=120, elem_id="ppp_combinatorial_limit", ) - return [force_equal_seeds, unlink_seed, seed, incremental_seed, combinatorial, combinatorial_limit] + return [force_equal_seeds, unlink_seed, seed, incremental_seed, combinatorial, combinatorial_shuffle, combinatorial_limit] def process( self, @@ -165,6 +169,7 @@ class PromptPostProcessorA1111Script(scripts.Script): input_seed, input_incremental_seed, input_combinatorial, + input_combinatorial_shuffle, input_combinatorial_limit, ): # pylint: disable=arguments-differ """ @@ -177,6 +182,7 @@ class PromptPostProcessorA1111Script(scripts.Script): 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_shuffle (bool): Flag indicating whether to shuffle the combinatorial results. input_combinatorial_limit (int): Maximum number of combinations (0 = no limit). Returns: @@ -241,6 +247,7 @@ class PromptPostProcessorA1111Script(scripts.Script): opts, "ppp_rem_removeextranetworktags", PromptPostProcessor.DEFAULT_CUP_REMOVE_EXTRANETWORK_TAGS ), do_combinatorial=input_combinatorial, + combinatorial_shuffle=input_combinatorial_shuffle, combinatorial_limit=max(num_seeds, int(input_combinatorial_limit)) if input_combinatorial else 0, ) if self.ppp_logger is None: @@ -392,19 +399,40 @@ class PromptPostProcessorA1111Script(scripts.Script): 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))] + 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))] + hiresfix_equal = regular_exists and rph == rpr and rnh == rnr + if hiresfix_equal: + log( + self.ppp_logger, + self.ppp_debug_level, + logging.INFO, + "hiresfix prompts are the same as regular prompts, skipping combinatorial processing for hiresfix", + ) + for i in range(len(rph)): # pylint: disable=consider-using-enumerate + prompts_list[(hiresfix_type, i)] = prompts_list.get((regular_type, i)) + else: + 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}])") + 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 @@ -412,7 +440,7 @@ class PromptPostProcessorA1111Script(scripts.Script): ) cached = self.lru_cache.get(key) if cached is None: - (hsh, seed, prompt, negative_prompt) = key + hsh, seed, prompt, negative_prompt = key results = ppp.process_prompt(prompt, negative_prompt, seed) posp, negp, _ = results[0] cached = (posp, negp) @@ -423,14 +451,14 @@ class PromptPostProcessorA1111Script(scripts.Script): 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: + # with open(os.path.join(os.path.dirname(os.path.realpath(__file__)), "..", "logs", f"last_prompts_{app.value}.txt"), "w", encoding="utf-8") as f: # for (prompttype, typeindex), (posp, negp) in prompts_list.items(): - # f.write(f"Key: {prompttype} {typeindex}\n") + # f.write(f"Key: {prompttype}[{typeindex}]\n") # f.write(f"Seed: {calculated_seeds[typeindex]}\n") - # f.write(f"Old Positive: {rpr[typeindex] if prompttype == regular_type else rph[typeindex]}\n") - # f.write(f"Old Negative: {rnr[typeindex] if prompttype == regular_type else rnh[typeindex]}\n") - # f.write(f"New Positive: {posp}\n") - # f.write(f"New Negative: {negp}\n") + # f.write(f"In Positive: {rpr[typeindex] if prompttype == regular_type else rph[typeindex]}\n") + # f.write(f"In Negative: {rnr[typeindex] if prompttype == regular_type else rnh[typeindex]}\n") + # f.write(f"Out Positive: {posp}\n") + # f.write(f"Out Negative: {negp}\n") # f.write("\n") # updates the prompts diff --git a/tests/base_tests.py b/tests/base_tests.py index cd9e93b..2b40e6d 100644 --- a/tests/base_tests.py +++ b/tests/base_tests.py @@ -1,7 +1,7 @@ from dataclasses import replace import os import logging -from typing import NamedTuple, Optional +from typing import Any, NamedTuple, Optional import unittest import datetime @@ -12,7 +12,7 @@ from ppp import PromptPostProcessor # type: ignore from ppp_logging import DEBUG_LEVEL, PromptPostProcessorLogFactory # type: ignore -class PromptPair(NamedTuple): +class InputTuple(NamedTuple): prompt: str = "" negative_prompt: str = "" @@ -20,7 +20,7 @@ class PromptPair(NamedTuple): class OutputTuple(NamedTuple): prompt: str = "" negative_prompt: str = "" - variables: dict[str, str] = None + variables: dict[str, Any] = None class TestPromptPostProcessorBase(unittest.TestCase): @@ -177,7 +177,7 @@ class TestPromptPostProcessorBase(unittest.TestCase): def process( self, - input_prompts: PromptPair, + input_prompts: InputTuple, expected_output: Optional[OutputTuple | list[OutputTuple]] = None, seed: int = 1, ppp: Optional[str | PromptPostProcessor] = None, @@ -188,7 +188,7 @@ class TestPromptPostProcessorBase(unittest.TestCase): Process the prompt and compare the results with the expected prompts. Args: - input_prompts (PromptPair): The input prompts. + input_prompts (InputTuple): 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. @@ -214,7 +214,7 @@ class TestPromptPostProcessorBase(unittest.TestCase): ) if self.interrupted != interrupted: errors.append(f"Interrupted flag is incorrect: expected {interrupted}, got {self.interrupted}") - elif expected_output is not None: + elif not self.interrupted and expected_output is not None: if len(result) != len(out): errors.append(f"Incorrect number of combinations (expected {len(out)}, got {len(result)})") for out_prompt, out_negative_prompt, out_variables in out: @@ -253,8 +253,8 @@ class TestPromptPostProcessorBase(unittest.TestCase): ) if self.interrupted != interrupted: errors.append(f"Interrupted flag is incorrect: expected {interrupted}, got {self.interrupted}") - elif expected_output is not None: - result_prompt, result_negative_prompt, output_variables = result[0] + elif not self.interrupted and expected_output is not None: + result_prompt, result_negative_prompt, output_variables = (result[0] if result else (None, None, None)) if result_prompt != eo.prompt or result_negative_prompt != eo.negative_prompt: errors.append( f"Incorrect result '{eo.prompt}' / '{eo.negative_prompt}', got '{result_prompt}' / '{result_negative_prompt}'" diff --git a/tests/tests_choices.py b/tests/tests_choices.py index f1fb35c..b2ed022 100644 --- a/tests/tests_choices.py +++ b/tests/tests_choices.py @@ -1,7 +1,7 @@ from dataclasses import replace from ppp import PromptPostProcessor # type: ignore -from .base_tests import OutputTuple, PromptPair, TestPromptPostProcessorBase +from .base_tests import OutputTuple, InputTuple, TestPromptPostProcessorBase if __name__ == "__main__": raise SystemExit("This script must not be run directly") @@ -16,14 +16,14 @@ class TestChoices(TestPromptPostProcessorBase): def test_ch_choices(self): # simple choices with weights self.process( - PromptPair("the choices are: {3::choice1|2::choice2|choice3}", ""), + InputTuple("the choices are: {3::choice1|2::choice2|choice3}", ""), OutputTuple("the choices are: choice2", ""), ppp="nocup", ) def test_ch_unsupportedsampler(self): # unsupported sampler self.process( - PromptPair("the choices are: {@choice1|choice2|choice3}", ""), + InputTuple("the choices are: {@choice1|choice2|choice3}", ""), OutputTuple("", ""), ppp="nocup", interrupted=True, @@ -31,7 +31,7 @@ class TestChoices(TestPromptPostProcessorBase): def test_ch_choices_withcomments(self): # choices with comments and multiline self.process( - PromptPair( + InputTuple( "the choices are: {\n3::choice1 # this is option 1\n|2::choice2\n# this was option 2\n|choice3 # this is option 3\n}", "", ), @@ -41,28 +41,28 @@ class TestChoices(TestPromptPostProcessorBase): def test_ch_choices_multiple(self): # choices with multiple selection self.process( - PromptPair("the choices are: {~2$$, $$3::choice1|2:: choice2 |choice3}", ""), + InputTuple("the choices are: {~2$$, $$3::choice1|2:: 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}", ""), + InputTuple("the choices are: {2$$, $$3::choice1|2 if _is_sd1::choice2|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}", ""), + InputTuple("${var=test}the choices are: {2$$, $$3::choice1|2 if not var eq 'test'::choice2|choice3}", ""), OutputTuple("the choices are: choice1, choice3", ""), ppp="nocup", ) def test_ch_choices_set_if_nested(self): # nested choices with if user variable and multiple selection self.process( - PromptPair( + InputTuple( "${var=test}the choices are: {2$$, $$3::choice1${var2=test2} {if var2 eq 'test2'::choice11|choice12}|2 if not var eq 'test'::choice2|choice3}", "", ), @@ -72,14 +72,14 @@ class TestChoices(TestPromptPostProcessorBase): def test_ch_choicesinsidelora(self): # simple choices inside a lora self.process( - PromptPair("", ""), + InputTuple("", ""), OutputTuple("", ""), ppp="nocup", ) def test_ch_removelorawithchoices(self): self.process( - PromptPair("", ""), + InputTuple("", ""), OutputTuple("", ""), ppp=PromptPostProcessor( self.ppp_logger, @@ -97,7 +97,7 @@ class TestChoices(TestPromptPostProcessorBase): def test_ch_cmd_includewildcard(self): self.process( - PromptPair("{ch_one|ch_two|%0.5::include yaml/wildcard1}", ""), + InputTuple("{ch_one|ch_two|%0.5::include yaml/wildcard1}", ""), OutputTuple("ch_two", ""), ppp="nocup", ) @@ -106,7 +106,7 @@ class TestChoices(TestPromptPostProcessorBase): def test_ch_combinatorial(self): self.process( - PromptPair("{choice1|choice2|choice3}, ${v:{option1|option2}}", ""), + InputTuple("{choice1|choice2|choice3}, ${v:{option1|option2}}", ""), [ OutputTuple("choice1, option1", ""), OutputTuple("choice1, option2", ""), diff --git a/tests/tests_cleanup.py b/tests/tests_cleanup.py index a05a7d0..b74e11c 100644 --- a/tests/tests_cleanup.py +++ b/tests/tests_cleanup.py @@ -2,7 +2,7 @@ import logging from dataclasses import replace from ppp import PromptPostProcessor # type: ignore -from .base_tests import OutputTuple, PromptPair, TestPromptPostProcessorBase +from .base_tests import OutputTuple, InputTuple, TestPromptPostProcessorBase if __name__ == "__main__": @@ -18,13 +18,13 @@ 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 "), + InputTuple(" this is a ((test ), , , (), , [] ( , test ,:2.0):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 self.process( - PromptPair( + InputTuple( " 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 ", ), @@ -36,7 +36,7 @@ class TestCleanup(TestPromptPostProcessorBase): def test_cl_removenetworktags(self): # remove network tags self.process( - PromptPair("this is a test__yaml/wildcard7__", ""), + InputTuple("this is a test__yaml/wildcard7__", ""), OutputTuple("this is a test", ""), ppp=PromptPostProcessor( self.ppp_logger, @@ -54,7 +54,7 @@ class TestCleanup(TestPromptPostProcessorBase): def test_cl_dontremoveseparatorsoneol(self): # don't remove separators on eol self.process( - PromptPair("this is a test,\nsecond line", ""), + InputTuple("this is a test,\nsecond line", ""), OutputTuple("this is a test,\nsecond line", ""), ppp=PromptPostProcessor( self.ppp_logger, @@ -73,7 +73,7 @@ class TestCleanup(TestPromptPostProcessorBase): def test_cl_separatorswitheol(self): # don't remove eols with the separators self.process( - PromptPair( + InputTuple( """{ (d:0.9) ,, (l:1.1) | (l:1.1) (d:0.9),,, } (l:1.1) (d:0.9)""", @@ -111,7 +111,7 @@ class TestCleanup(TestPromptPostProcessorBase): def test_cl_mergeattention(self): # merge attention self.process( - PromptPair( + InputTuple( "this is (a test:0.9) of (attention (merging:1.2)) where ((this)) ((is joined:1.2)) and ([this too]:1.3)", "", ), @@ -123,7 +123,7 @@ class TestCleanup(TestPromptPostProcessorBase): def test_cl_not_mergeattention(self): # not merge attention self.process( - PromptPair( + InputTuple( "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)", "", ), @@ -139,7 +139,7 @@ class TestCleanup(TestPromptPostProcessorBase): def test_cl_warn_unmatched_open_paren(self): # unmatched open parenthesis triggers warning with self.assertLogs("PromptPostProcessor", level=logging.WARNING) as cm: self.process( - PromptPair("(unclosed paren", ""), + InputTuple("(unclosed paren", ""), OutputTuple("(unclosed paren", ""), ) self.assertTrue( @@ -150,7 +150,7 @@ class TestCleanup(TestPromptPostProcessorBase): def test_cl_warn_unmatched_close_paren(self): # unmatched close parenthesis triggers warning with self.assertLogs("PromptPostProcessor", level=logging.WARNING) as cm: self.process( - PromptPair("extra close paren)", ""), + InputTuple("extra close paren)", ""), OutputTuple("extra close paren)", ""), ) self.assertTrue( @@ -161,7 +161,7 @@ class TestCleanup(TestPromptPostProcessorBase): def test_cl_warn_mismatched_brackets(self): # mismatched bracket types trigger warning with self.assertLogs("PromptPostProcessor", level=logging.WARNING) as cm: self.process( - PromptPair("(mismatched]", ""), + InputTuple("(mismatched]", ""), OutputTuple("(mismatched]", ""), ) self.assertTrue( @@ -172,7 +172,7 @@ class TestCleanup(TestPromptPostProcessorBase): def test_cl_warn_unmatched_open_bracket(self): # unmatched open bracket triggers warning with self.assertLogs("PromptPostProcessor", level=logging.WARNING) as cm: self.process( - PromptPair("unclosed [bracket", ""), + InputTuple("unclosed [bracket", ""), OutputTuple("unclosed [bracket", ""), ) self.assertTrue( @@ -183,7 +183,7 @@ class TestCleanup(TestPromptPostProcessorBase): def test_cl_warn_unmatched_complex(self): # unmatched complex case triggers warning with self.assertLogs("PromptPostProcessor", level=logging.WARNING) as cm: self.process( - PromptPair("[(unmatched [bracket))", ""), + InputTuple("[(unmatched [bracket))", ""), OutputTuple("[(unmatched [bracket))", ""), ) self.assertTrue( @@ -194,6 +194,6 @@ class TestCleanup(TestPromptPostProcessorBase): def test_cl_warn_escaped_unmatched_no_false_warning(self): # escaped unmatched paren/bracket does not trigger warning with self.assertNoLogs("PromptPostProcessor", level=logging.WARNING): self.process( - PromptPair(r"text with \(escaped unmatched\]", ""), + InputTuple(r"text with \(escaped unmatched\]", ""), OutputTuple(r"text with \(escaped unmatched\]", ""), ) diff --git a/tests/tests_host.py b/tests/tests_host.py index b58f6be..033931d 100644 --- a/tests/tests_host.py +++ b/tests/tests_host.py @@ -1,5 +1,5 @@ from ppp import PromptPostProcessor # type: ignore -from .base_tests import OutputTuple, PromptPair, TestPromptPostProcessorBase +from .base_tests import OutputTuple, InputTuple, TestPromptPostProcessorBase if __name__ == "__main__": raise SystemExit("This script must not be run directly") @@ -14,7 +14,7 @@ class TestHosts(TestPromptPostProcessorBase): def test_host_attention_parentheses(self): self.process( - PromptPair( + InputTuple( "[test1] (test2) (test3:1.5) [(test4)]", "", ), @@ -35,7 +35,7 @@ class TestHosts(TestPromptPostProcessorBase): def test_host_attention_disable(self): self.process( - PromptPair( + InputTuple( "[test1] (test2) (test3:1.5)", "", ), @@ -56,7 +56,7 @@ class TestHosts(TestPromptPostProcessorBase): def test_host_attention_remove(self): self.process( - PromptPair( + InputTuple( "[test1] (test2) (test3:1.5)", "", ), @@ -77,7 +77,7 @@ class TestHosts(TestPromptPostProcessorBase): def test_host_attention_error(self): self.process( - PromptPair( + InputTuple( "[test1] (test2) (test3:1.5)", "", ), @@ -99,7 +99,7 @@ class TestHosts(TestPromptPostProcessorBase): def test_host_scheduling_before(self): self.process( - PromptPair( + InputTuple( "[test1:test2:0.5]", "", ), @@ -120,7 +120,7 @@ class TestHosts(TestPromptPostProcessorBase): def test_host_scheduling_after(self): self.process( - PromptPair( + InputTuple( "[test1:test2:0.5]", "", ), @@ -141,7 +141,7 @@ class TestHosts(TestPromptPostProcessorBase): def test_host_scheduling_first(self): self.process( - PromptPair( + InputTuple( "[test1::0.5] [:test2:0.5] [test3:test4:0.5]", "", ), @@ -162,7 +162,7 @@ class TestHosts(TestPromptPostProcessorBase): def test_host_scheduling_remove(self): self.process( - PromptPair( + InputTuple( "[test1:test2:0.5]", "", ), @@ -183,7 +183,7 @@ class TestHosts(TestPromptPostProcessorBase): def test_host_scheduling_error(self): self.process( - PromptPair( + InputTuple( "[test1:test2:0.5]", "", ), @@ -205,7 +205,7 @@ class TestHosts(TestPromptPostProcessorBase): def test_host_alternation_first(self): self.process( - PromptPair( + InputTuple( "[test1|test2|test3]", "", ), @@ -226,7 +226,7 @@ class TestHosts(TestPromptPostProcessorBase): def test_host_alternation_remove(self): self.process( - PromptPair( + InputTuple( "[test1|test2|test3]", "", ), @@ -247,7 +247,7 @@ class TestHosts(TestPromptPostProcessorBase): def test_host_alternation_error(self): self.process( - PromptPair( + InputTuple( "[test1|test2|test3]", "", ), @@ -269,7 +269,7 @@ class TestHosts(TestPromptPostProcessorBase): def test_host_and_eol(self): self.process( - PromptPair( + InputTuple( "test1 AND test2:2", "", ), @@ -290,7 +290,7 @@ class TestHosts(TestPromptPostProcessorBase): def test_host_and_comma(self): self.process( - PromptPair( + InputTuple( "test1 AND test2:2", "", ), @@ -311,7 +311,7 @@ class TestHosts(TestPromptPostProcessorBase): def test_host_and_remove(self): self.process( - PromptPair( + InputTuple( "test1 AND test2:2", "", ), @@ -332,7 +332,7 @@ class TestHosts(TestPromptPostProcessorBase): def test_host_and_error(self): self.process( - PromptPair( + InputTuple( "test1 AND test2:2", "", ), @@ -354,7 +354,7 @@ class TestHosts(TestPromptPostProcessorBase): def test_host_break_eol(self): self.process( - PromptPair( + InputTuple( "test1 BREAK test2", "", ), @@ -375,7 +375,7 @@ class TestHosts(TestPromptPostProcessorBase): def test_host_break_comma(self): self.process( - PromptPair( + InputTuple( "test1 BREAK test2", "", ), @@ -396,7 +396,7 @@ class TestHosts(TestPromptPostProcessorBase): def test_host_break_remove(self): self.process( - PromptPair( + InputTuple( "test1 BREAK test2", "", ), @@ -417,7 +417,7 @@ class TestHosts(TestPromptPostProcessorBase): def test_host_break_error(self): self.process( - PromptPair( + InputTuple( "test1 BREAK test2", "", ), diff --git a/tests/tests_performance.py b/tests/tests_performance.py index 668b643..a5be804 100644 --- a/tests/tests_performance.py +++ b/tests/tests_performance.py @@ -1,4 +1,4 @@ -from .base_tests import PromptPair, TestPromptPostProcessorBase +from .base_tests import InputTuple, TestPromptPostProcessorBase if __name__ == "__main__": raise SystemExit("This script must not be run directly") @@ -18,7 +18,7 @@ class TestPerformance(TestPromptPostProcessorBase): ["(this:1.2) is a [test] using a [simple|low complexity] prompt with "] * 15 ) self.process( - PromptPair(large_prompt, ""), + InputTuple(large_prompt, ""), ppp="nocup", ) @@ -30,7 +30,7 @@ class TestPerformance(TestPromptPostProcessorBase): ["(this:1.2) is a [test] using a [simple|low complexity] prompt with "] * 15 ) self.process( - PromptPair(large_prompt, ""), + InputTuple(large_prompt, ""), ppp="nocup", ) @@ -39,7 +39,7 @@ class TestPerformance(TestPromptPostProcessorBase): ): # performance test with a large prompt with new constructs (full parser) large_prompt = ", ".join(["__yaml/wildcard1__, (__yaml/wildcard2__), __yaml/wildcard3__, {one|two|three}"] * 15) self.process( - PromptPair(large_prompt, ""), + InputTuple(large_prompt, ""), ppp="nocup", ) @@ -49,27 +49,27 @@ class TestPerformance(TestPromptPostProcessorBase): def test_parser_performance_simple_attention(self): # performance test with only attention large_prompt = ", ".join(["(one:1.2) two (three) four [five] six"] * 20) self.process( - PromptPair(large_prompt, ""), + InputTuple(large_prompt, ""), ppp="nocup", ) def test_parser_performance_simple_schedules(self): # performance test with only schedules large_prompt = ", ".join(["[one:1:0.5] two [three:0.8] four [five:5:0.2] six"] * 20) self.process( - PromptPair(large_prompt, ""), + InputTuple(large_prompt, ""), ppp="nocup", ) def test_parser_performance_simple_alternation(self): # performance test with only alternation large_prompt = ", ".join(["[one|1] two [three|3] four [five|5] six"] * 20) self.process( - PromptPair(large_prompt, ""), + InputTuple(large_prompt, ""), ppp="nocup", ) def test_parser_performance_simple_extranetwork(self): # performance test with only extra networks large_prompt = ", ".join([" two four six"] * 20) self.process( - PromptPair(large_prompt, ""), + InputTuple(large_prompt, ""), ppp="nocup", ) diff --git a/tests/tests_stn.py b/tests/tests_stn.py index 5cc2e40..50fe1a5 100644 --- a/tests/tests_stn.py +++ b/tests/tests_stn.py @@ -1,4 +1,4 @@ -from .base_tests import OutputTuple, PromptPair, TestPromptPostProcessorBase +from .base_tests import OutputTuple, InputTuple, TestPromptPostProcessorBase if __name__ == "__main__": raise SystemExit("This script must not be run directly") @@ -13,7 +13,7 @@ class TestSendToNegative(TestPromptPostProcessorBase): def test_stn_simple(self): # negtags with different parameters and separations self.process( - PromptPair( + InputTuple( "flowersred, green, blueyellow, purpleblack", "normal quality, worse quality", ), @@ -22,7 +22,7 @@ class TestSendToNegative(TestPromptPostProcessorBase): def test_stn_complex(self): # complex negtags self.process( - PromptPair( + InputTuple( "red ((pink)), flowers purple, mauveblue, yellow green", "normal quality, , bad quality, worse quality", ), @@ -34,7 +34,7 @@ class TestSendToNegative(TestPromptPostProcessorBase): def test_stn_complex_nocleanup(self): # complex negtags with no cleanup self.process( - PromptPair( + InputTuple( "red ((pink)), flowers purple, mauveblue, yellow green", "normal quality, , bad quality, worse quality", ), @@ -47,7 +47,7 @@ class TestSendToNegative(TestPromptPostProcessorBase): def test_stn_inside_attention(self): # negtag inside attention self.process( - PromptPair( + InputTuple( "[neg1] this is a ((testneg2) (test:2.0): 1.5 ) (red[square]:1.5)", "normal quality", ), @@ -59,7 +59,7 @@ class TestSendToNegative(TestPromptPostProcessorBase): def test_stn_inside_alternation(self): # negtag inside alternation self.process( - PromptPair( + InputTuple( "this is a (([complexneg1|simpleneg2|regularneg3] test)(test:2.0):1.5)", "normal quality", ), @@ -71,7 +71,7 @@ class TestSendToNegative(TestPromptPostProcessorBase): def test_stn_inside_alternation_recursive(self): # negtag inside alternation (recursive alternation) self.process( - PromptPair( + InputTuple( "this is a (([complexneg1[one|twoneg12||three|four(neg14)]|simpleneg2|regularneg3] test)(test:2.0):1.5)", "normal quality", ), @@ -83,13 +83,13 @@ class TestSendToNegative(TestPromptPostProcessorBase): def test_stn_inside_scheduling(self): # negtag inside scheduling self.process( - PromptPair("this is [abcneg1:defneg2: 5 ]", "normal quality"), + InputTuple("this is [abcneg1:defneg2: 5 ]", "normal quality"), 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 self.process( - PromptPair( + InputTuple( "[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, ", ), @@ -101,7 +101,7 @@ class TestSendToNegative(TestPromptPostProcessorBase): def test_stn_complex_features_newformat(self): # complex negtags with AND, BREAK and other features (new format) self.process( - PromptPair( + InputTuple( "[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, ", ), @@ -113,7 +113,7 @@ class TestSendToNegative(TestPromptPostProcessorBase): def test_stn_inside_alternation_recursive_2(self): # negtag inside alternation (recursive alternation) self.process( - PromptPair( + InputTuple( "[pos1neg1[pos11|pos12neg12||pos14|pos15neg15]|pos2neg2|pos3neg3]", "", ), diff --git a/tests/tests_varcomms.py b/tests/tests_varcomms.py index d85bb3e..d5139f4 100644 --- a/tests/tests_varcomms.py +++ b/tests/tests_varcomms.py @@ -2,7 +2,7 @@ from dataclasses import replace from ppp import PromptPostProcessor # type: ignore from ppp_classes import ONWARNING_CHOICES # type: ignore -from .base_tests import OutputTuple, PromptPair, TestPromptPostProcessorBase +from .base_tests import OutputTuple, InputTuple, TestPromptPostProcessorBase if __name__ == "__main__": raise SystemExit("This script must not be run directly") @@ -17,7 +17,7 @@ class TestVarCommands(TestPromptPostProcessorBase): def test_empty_variable(self): self.process( - PromptPair( + InputTuple( "${v1=}${v3:}", "", ), @@ -28,7 +28,7 @@ class TestVarCommands(TestPromptPostProcessorBase): def test_echoed_variable(self): self.process( - PromptPair( + InputTuple( "${v1=test1}test2${v3:test3}${v3:test4}", "", ), @@ -38,7 +38,7 @@ class TestVarCommands(TestPromptPostProcessorBase): def test_unknown_echoed_variable(self): self.process( - PromptPair( + InputTuple( "${v1}", "", ), @@ -61,7 +61,7 @@ class TestVarCommands(TestPromptPostProcessorBase): def test_variable_in_extranetwork(self): self.process( - PromptPair( + InputTuple( "${f=filename}${w=0.5}", "", ), @@ -72,7 +72,7 @@ class TestVarCommands(TestPromptPostProcessorBase): def test_var_nested_1(self): # variable default nested in variable set self.process( - PromptPair( + InputTuple( "${v1=test ${v2:OK}}${v1}", "", ), @@ -81,7 +81,7 @@ class TestVarCommands(TestPromptPostProcessorBase): def test_var_nested_2(self): # variable set nested in variable default self.process( - PromptPair( + InputTuple( "${v1:test ${v2=OK}${v2}}", "", ), @@ -90,7 +90,7 @@ class TestVarCommands(TestPromptPostProcessorBase): def test_var_nested_3(self): # variable default nested in variable default self.process( - PromptPair( + InputTuple( "${v1:test ${v2:OK}}", "", ), @@ -101,7 +101,7 @@ class TestVarCommands(TestPromptPostProcessorBase): def test_array_variable_1(self): # array variable set with += and test of index value and full array with and without default separator self.process( - PromptPair( + InputTuple( "${v1[]=val1}${v1[]+=val2}${v1[]+=val3}${v1[1]:defval},${v1[]:defval2},${v1[&'.']:defval3}", "", ), @@ -110,7 +110,7 @@ class TestVarCommands(TestPromptPostProcessorBase): 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 self.process( - PromptPair( + InputTuple( "${v1[]=val1}${v1[]=val2}${v1[]:defval},${v2[]:defval2},${v2[1]:defval3},${v3[]=}${v3[]:defval4}", "", ), @@ -119,7 +119,7 @@ class TestVarCommands(TestPromptPostProcessorBase): def test_array_variable_3(self): # access array index by variable, set array variable to expanded array variable and add expanded array self.process( - PromptPair( + InputTuple( "${v1[]=val1}${v1[]+=val2}${v2=1}${v1[v2]:defval1}${v3[]=${v1[]}}${v3[]+=${v1[]}}, ${v3[&'.']}", "", ), @@ -128,7 +128,7 @@ class TestVarCommands(TestPromptPostProcessorBase): def test_array_variable_4(self): # test list in array self.process( - PromptPair( + InputTuple( "${v1[]=val1}${v1[]+=val2}${v1[]+=val3}OKnot OK", "", ), @@ -137,7 +137,7 @@ class TestVarCommands(TestPromptPostProcessorBase): def test_array_variable_5(self): # test empty array self.process( - PromptPair( + InputTuple( "${v1[]=}OKnot OK,OKnot OK", "", ), @@ -146,7 +146,7 @@ class TestVarCommands(TestPromptPostProcessorBase): def test_array_variable_6(self): # array variable set and addition with expanded values from array variables self.process( - PromptPair( + InputTuple( "${v1[]=val1}${v1[]+=val2}${v2[]=val3}${v3[]=*v1[]}${v3[]+=*v2[]}", "", ), @@ -155,7 +155,7 @@ class TestVarCommands(TestPromptPostProcessorBase): def test_array_variable_7(self): # array variable set and addition with expanded values from wildcards self.process( - PromptPair( + InputTuple( "${v1[]=*__yaml/wildcard1__}${v1[]+=*__yaml/wildcard2__}${v1[2]:defval}", "", ), @@ -164,7 +164,7 @@ class TestVarCommands(TestPromptPostProcessorBase): def test_array_variable_8(self): # array variable set and addition with expanded values from lists self.process( - PromptPair( + InputTuple( "${v1[]=*()}${v1[]+=*('one','two')}${v2=three}${v1[]+=*(v2,'four')}${v1[2]:defval}", "", ), @@ -173,7 +173,7 @@ class TestVarCommands(TestPromptPostProcessorBase): def test_array_variable_9(self): # array variable length self.process( - PromptPair( + InputTuple( "${v1[]=val1}${v1[]+=val2}${v1[]+=val3}${v1[#]:defval}, OKnot OK", "", ), @@ -182,7 +182,7 @@ class TestVarCommands(TestPromptPostProcessorBase): def test_array_variable_10(self): # array variable set with expanded values from wildcards in command format self.process( - PromptPair( + InputTuple( "*__yaml/wildcard1__*__yaml/wildcard2__${v1[2]:defval}", "", ), @@ -195,7 +195,7 @@ class TestVarCommands(TestPromptPostProcessorBase): def test_operator_ReqR(self): self.process( - PromptPair( + InputTuple( "${r1=hello}${r2=hello}OKnot OK", "", ), @@ -204,7 +204,7 @@ class TestVarCommands(TestPromptPostProcessorBase): def test_operator_RnoteqR(self): # test for the not before the operator self.process( - PromptPair( + InputTuple( "${r1=hello}${r2=bye}OKnot OK", "", ), @@ -213,7 +213,7 @@ class TestVarCommands(TestPromptPostProcessorBase): def test_operator_RneR(self): self.process( - PromptPair( + InputTuple( "${r1=hello}${r2=bye}OKnot OK", "", ), @@ -222,7 +222,7 @@ class TestVarCommands(TestPromptPostProcessorBase): def test_operator_RltR(self): self.process( - PromptPair( + InputTuple( "${r1=1}${r2=2}OKnot OK", "", ), @@ -231,7 +231,7 @@ class TestVarCommands(TestPromptPostProcessorBase): def test_operator_RgtR(self): self.process( - PromptPair( + InputTuple( "${r1=2}${r2=1}OKnot OK", "", ), @@ -240,7 +240,7 @@ class TestVarCommands(TestPromptPostProcessorBase): def test_operator_RleR(self): self.process( - PromptPair( + InputTuple( "${r1=1}${r2=1}OKnot OK", "", ), @@ -249,7 +249,7 @@ class TestVarCommands(TestPromptPostProcessorBase): def test_operator_RgeR(self): self.process( - PromptPair( + InputTuple( "${r1=1}${r2=1}OKnot OK", "", ), @@ -258,7 +258,7 @@ class TestVarCommands(TestPromptPostProcessorBase): def test_operator_RinR(self): self.process( - PromptPair( + InputTuple( "${r1=hello}${r2=hello world}OKnot OK", "", ), @@ -267,7 +267,7 @@ class TestVarCommands(TestPromptPostProcessorBase): def test_operator_RcontainsR(self): self.process( - PromptPair( + InputTuple( "${r1=hello world}${r2=hello}OKnot OK", "", ), @@ -278,7 +278,7 @@ class TestVarCommands(TestPromptPostProcessorBase): def test_operator_AeqA(self): self.process( - PromptPair( + InputTuple( "${a1[]=*('hello','world')}${a2[]=*('hello','world')}OKnot OK", "", ), @@ -287,7 +287,7 @@ class TestVarCommands(TestPromptPostProcessorBase): def test_operator_AneA_1(self): self.process( - PromptPair( + InputTuple( "${a1[]=*('hello')}${a2[]=*('bye')}OKnot OK", "", ), @@ -296,7 +296,7 @@ class TestVarCommands(TestPromptPostProcessorBase): def test_operator_AneA_2(self): self.process( - PromptPair( + InputTuple( "${a1[]=*('hello','world')}${a2[]=*('hello')}OKnot OK", "", ), @@ -305,7 +305,7 @@ class TestVarCommands(TestPromptPostProcessorBase): def test_operator_AneA_3(self): self.process( - PromptPair( + InputTuple( "${a1[]=*('hello','world')}${a2[]=*('world','hello')}OKnot OK", "", ), @@ -314,7 +314,7 @@ class TestVarCommands(TestPromptPostProcessorBase): def test_operator_AltA_1(self): self.process( - PromptPair( + InputTuple( "${a1[]=*(1,2,3)}${a2[]=*(2,3,4)}OKnot OK", "", ), @@ -323,7 +323,7 @@ class TestVarCommands(TestPromptPostProcessorBase): def test_operator_AltA_2(self): self.process( - PromptPair( + InputTuple( "${a1[]=*(1,2)}${a2[]=*(2,3,4)}OKnot OK", "", ), @@ -332,7 +332,7 @@ class TestVarCommands(TestPromptPostProcessorBase): def test_operator_AgtA(self): self.process( - PromptPair( + InputTuple( "${a1[]=*(2,3,4)}${a2[]=*(1,2,3)}OKnot OK", "", ), @@ -341,7 +341,7 @@ class TestVarCommands(TestPromptPostProcessorBase): def test_operator_AleA(self): self.process( - PromptPair( + InputTuple( "${a1[]=*(1,2)}${a2[]=*(1,3)}OKnot OK", "", ), @@ -350,7 +350,7 @@ class TestVarCommands(TestPromptPostProcessorBase): def test_operator_AgeA(self): self.process( - PromptPair( + InputTuple( "${a1[]=*(1,3)}${a2[]=*(1,2)}OKnot OK", "", ), @@ -359,7 +359,7 @@ class TestVarCommands(TestPromptPostProcessorBase): def test_operator_AinA(self): self.process( - PromptPair( + InputTuple( "${a1[]=*('hello')}${a2[]=*('hello', 'world')}OKnot OK", "", ), @@ -368,7 +368,7 @@ class TestVarCommands(TestPromptPostProcessorBase): def test_operator_AcontainsA(self): self.process( - PromptPair( + InputTuple( "${a1[]=*('hello','world')}${a2[]=*('hello')}OKnot OK", "", ), @@ -379,7 +379,7 @@ class TestVarCommands(TestPromptPostProcessorBase): def test_operator_AeqR(self): self.process( - PromptPair( + InputTuple( "${a1[]=*('hello','world')}${r2=hello}OKnot OK", "", ), @@ -389,7 +389,7 @@ class TestVarCommands(TestPromptPostProcessorBase): def test_operator_AneR(self): self.process( - PromptPair( + InputTuple( "${a1[]=*('hello')}${r2=bye}OKnot OK", "", ), @@ -399,7 +399,7 @@ class TestVarCommands(TestPromptPostProcessorBase): def test_operator_AltR(self): self.process( - PromptPair( + InputTuple( "${a1[]=*(1,2,3)}${r2=2}OKnot OK", "", ), @@ -409,7 +409,7 @@ class TestVarCommands(TestPromptPostProcessorBase): def test_operator_AgtR(self): self.process( - PromptPair( + InputTuple( "${a1[]=*(2,3,4)}${r2=2}OKnot OK", "", ), @@ -419,7 +419,7 @@ class TestVarCommands(TestPromptPostProcessorBase): def test_operator_AleR(self): self.process( - PromptPair( + InputTuple( "${a1[]=*(1,2)}${r2=2}OKnot OK", "", ), @@ -429,7 +429,7 @@ class TestVarCommands(TestPromptPostProcessorBase): def test_operator_AgeR(self): self.process( - PromptPair( + InputTuple( "${a1[]=*(1,3)}${r2=2}OKnot OK", "", ), @@ -439,7 +439,7 @@ class TestVarCommands(TestPromptPostProcessorBase): def test_operator_AinR(self): self.process( - PromptPair( + InputTuple( "${a1[]=*('hello', 'world')}${r2=hello world)}OKnot OK", "", ), @@ -448,7 +448,7 @@ class TestVarCommands(TestPromptPostProcessorBase): def test_operator_AcontainsR(self): self.process( - PromptPair( + InputTuple( "${a1[]=*('hello','world')}${r2=hello}OKnot OK", "", ), @@ -459,7 +459,7 @@ class TestVarCommands(TestPromptPostProcessorBase): def test_operator_ReqA(self): self.process( - PromptPair( + InputTuple( "${r1=hello}${a2[]=*('hello','world')}OKnot OK", "", ), @@ -469,7 +469,7 @@ class TestVarCommands(TestPromptPostProcessorBase): def test_operator_RneA(self): self.process( - PromptPair( + InputTuple( "${r1=bye}${a2[]=*('hello')}OKnot OK", "", ), @@ -479,7 +479,7 @@ class TestVarCommands(TestPromptPostProcessorBase): def test_operator_RltA(self): self.process( - PromptPair( + InputTuple( "${r1=2}${a2[]=*(1,2,3)}OKnot OK", "", ), @@ -489,7 +489,7 @@ class TestVarCommands(TestPromptPostProcessorBase): def test_operator_RgtA(self): self.process( - PromptPair( + InputTuple( "${r1=2}${a2[]=*(2,3,4)}OKnot OK", "", ), @@ -499,7 +499,7 @@ class TestVarCommands(TestPromptPostProcessorBase): def test_operator_RleA(self): self.process( - PromptPair( + InputTuple( "${r1=2}${a2[]=*(1,2)}OKnot OK", "", ), @@ -509,7 +509,7 @@ class TestVarCommands(TestPromptPostProcessorBase): def test_operator_RgeA(self): self.process( - PromptPair( + InputTuple( "${r1=2}${a2[]=*(1,3)}OKnot OK", "", ), @@ -519,7 +519,7 @@ class TestVarCommands(TestPromptPostProcessorBase): def test_operator_RinA(self): self.process( - PromptPair( + InputTuple( "${r1=hello}${a2[]=*('hello', 'world')}OKnot OK", "", ), @@ -528,7 +528,7 @@ class TestVarCommands(TestPromptPostProcessorBase): def test_operator_RcontainsA(self): self.process( - PromptPair( + InputTuple( "${r1=hello world}${a2[]=*('hello','world')}OKnot OK", "", ), @@ -539,7 +539,7 @@ class TestVarCommands(TestPromptPostProcessorBase): def test_operator_ReqV_str(self): self.process( - PromptPair( + InputTuple( "${r1=hello}OKnot OK", "", ), @@ -548,7 +548,7 @@ class TestVarCommands(TestPromptPostProcessorBase): def test_operator_ReqV_str_fail(self): self.process( - PromptPair( + InputTuple( "${r1=hello}OKnot OK", "", ), @@ -558,7 +558,7 @@ class TestVarCommands(TestPromptPostProcessorBase): def test_operator_ReqV_num(self): self.process( - PromptPair( + InputTuple( "${r1=42}OKnot OK", "", ), @@ -567,7 +567,7 @@ class TestVarCommands(TestPromptPostProcessorBase): def test_operator_ReqV_num_fail(self): self.process( - PromptPair( + InputTuple( "${r1=42}OKnot OK", "", ), @@ -577,7 +577,7 @@ class TestVarCommands(TestPromptPostProcessorBase): def test_operator_ReqV_bool(self): self.process( - PromptPair( + InputTuple( "${r1=true}OKnot OK", "", ), @@ -586,7 +586,7 @@ class TestVarCommands(TestPromptPostProcessorBase): def test_operator_ReqV_bool_fail(self): self.process( - PromptPair( + InputTuple( "${r1=true}OKnot OK", "", ), @@ -598,7 +598,7 @@ class TestVarCommands(TestPromptPostProcessorBase): def test_listoperand_AinL(self): self.process( - PromptPair( + InputTuple( "${a1[]=*('hello','world')}OKnot OK", "", ), @@ -607,7 +607,7 @@ class TestVarCommands(TestPromptPostProcessorBase): def test_listoperand_LinA(self): self.process( - PromptPair( + InputTuple( "${a2[]=*('hello','world')}OKnot OK", "", ), @@ -618,7 +618,7 @@ class TestVarCommands(TestPromptPostProcessorBase): def test_indexedoperand_RinA(self): self.process( - PromptPair( + InputTuple( "${a2[]=*('hello','world')}OKnot OK", "", ), @@ -629,7 +629,7 @@ class TestVarCommands(TestPromptPostProcessorBase): def test_float_value(self): self.process( - PromptPair( + InputTuple( "${a=1.5}OKnot OK", "", ), @@ -641,7 +641,7 @@ class TestVarCommands(TestPromptPostProcessorBase): def test_cmd_if_undefined_var_int_compare_warn(self): # undefined var integer compare with on_warning=warn self.process( - PromptPair( + InputTuple( "OKnot OK", "", ), @@ -662,7 +662,7 @@ class TestVarCommands(TestPromptPostProcessorBase): def test_cmd_if_undefined_var_int_compare_stop(self): # undefined var integer compare with on_warning=stop self.process( - PromptPair( + InputTuple( "OKnot OK", "", ), @@ -672,7 +672,7 @@ class TestVarCommands(TestPromptPostProcessorBase): def test_cmd_if_nonnumeric_var_int_compare_warn(self): # non-numeric var integer compare with on_warning=warn self.process( - PromptPair( + InputTuple( "abcOKnot OK", "", ), @@ -693,7 +693,7 @@ class TestVarCommands(TestPromptPostProcessorBase): def test_cmd_if_nonnumeric_var_int_compare_stop(self): # non-numeric var integer compare with on_warning=stop self.process( - PromptPair( + InputTuple( "abcOKnot OK", "", ), @@ -703,7 +703,7 @@ class TestVarCommands(TestPromptPostProcessorBase): def test_cmd_if_empty_var_int_compare(self): # empty string var integer compare with on_warning=warn self.process( - PromptPair( + InputTuple( "OKnot OK", "", ), @@ -726,7 +726,7 @@ class TestVarCommands(TestPromptPostProcessorBase): def test_cmd_stn_complex_features(self): # complex stn command with AND, BREAK and other features self.process( - PromptPair( + InputTuple( "[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, ", ), @@ -738,7 +738,7 @@ class TestVarCommands(TestPromptPostProcessorBase): def test_cmd_if_complex_features(self): # complex if command self.process( - PromptPair( + InputTuple( "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", ), @@ -750,7 +750,7 @@ class TestVarCommands(TestPromptPostProcessorBase): def test_cmd_if_nested(self): # nested if command self.process( - PromptPair( + InputTuple( "this is SD1PONYSD2NOPONYNOPONY", "", ), @@ -771,25 +771,25 @@ class TestVarCommands(TestPromptPostProcessorBase): def test_cmd_set_if(self): # set and if commands self.process( - PromptPair("valuethis test is OKnot OK", ""), + InputTuple("valuethis test is OKnot OK", ""), OutputTuple("this test is OK", ""), ) def test_cmd_set_empty(self): # set to empty self.process( - PromptPair("${v2=}this test is not OKOK", ""), + InputTuple("${v2=}this test is not OKOK", ""), OutputTuple("this test is OK", ""), ) def test_cmd_set_eval_if(self): # set and if commands self.process( - PromptPair("valuethis test is OKnot OK", ""), + InputTuple("valuethis test is OKnot OK", ""), OutputTuple("this test is OK", ""), ) def test_cmd_set_if_echo_nested(self): # nested set, if and echo commands self.process( - PromptPair( + InputTuple( "1OKnot OK NOK OK", "", ), @@ -798,7 +798,7 @@ class TestVarCommands(TestPromptPostProcessorBase): def test_cmd_set_if_complex_conditions_1(self): # complex conditions (or) self.process( - PromptPair( + InputTuple( "truefalsethis test is OKnot OK", "", ), @@ -807,7 +807,7 @@ class TestVarCommands(TestPromptPostProcessorBase): def test_cmd_set_if_complex_conditions_2(self): # complex conditions (and) self.process( - PromptPair( + InputTuple( "truetruethis test is OKnot OK", "", ), @@ -816,13 +816,13 @@ class TestVarCommands(TestPromptPostProcessorBase): def test_cmd_set_if_complex_conditions_3(self): # complex conditions (not) self.process( - PromptPair("falsethis test is OKnot OK", ""), + InputTuple("falsethis test is OKnot OK", ""), OutputTuple("this test is OK", ""), ) def test_cmd_set_if_complex_conditions_4(self): # complex conditions (not, precedence) self.process( - PromptPair( + InputTuple( "truefalsethis test is OKnot OK", "", ), @@ -831,7 +831,7 @@ class TestVarCommands(TestPromptPostProcessorBase): def test_cmd_set_if_complex_conditions_5(self): # complex conditions (not, precedence, comparison) self.process( - PromptPair( + InputTuple( "1falsethis test is OKnot OK", "", ), @@ -840,7 +840,7 @@ class TestVarCommands(TestPromptPostProcessorBase): def test_cmd_set_if_complex_conditions_6(self): # complex conditions self.process( - PromptPair( + InputTuple( "123this test is OKnot OK", "", ), @@ -849,7 +849,7 @@ class TestVarCommands(TestPromptPostProcessorBase): def test_cmd_set_if_complex_conditions_7(self): # complex conditions self.process( - PromptPair( + InputTuple( "123this test is OKnot OK", "", ), @@ -858,7 +858,7 @@ class TestVarCommands(TestPromptPostProcessorBase): def test_cmd_set_if2(self): # set and more complex if commands self.process( - PromptPair( + InputTuple( "First: value1this test is OKOK2not OK\nSecond: value3this test is OKnot OK", "", ), @@ -867,7 +867,7 @@ class TestVarCommands(TestPromptPostProcessorBase): def test_cmd_set_add_if(self): # set, add and if commands self.process( - PromptPair( + InputTuple( "value2this test is OKnot OK", "", ), @@ -876,7 +876,7 @@ class TestVarCommands(TestPromptPostProcessorBase): def test_cmd_set_add_DP_if(self): # set, add (DP format) and if commands self.process( - PromptPair( + InputTuple( "${v=value}${v+=2}this test is OKnot OK", "", ), @@ -885,7 +885,7 @@ class TestVarCommands(TestPromptPostProcessorBase): def test_cmd_set_immediateeval(self): # set (DP format) with mixed evaluation self.process( - PromptPair( + InputTuple( "${var=!__yaml/wildcard1__}the choices are: ${var}, ${var}, ${var2:default}, ${var3=__yaml/wildcard1__}${var3}, ${var3}", "", ), @@ -895,7 +895,7 @@ class TestVarCommands(TestPromptPostProcessorBase): def test_cmd_set_mixeval(self): # set and add (DP format) with mixed evaluation self.process( - PromptPair( + InputTuple( "${var=__yaml/wildcard1__}the choices are: ${var}, ${var}, ${var+=, __yaml/wildcard2__}${var}, ${var}, ${var+=!, __yaml/wildcard3__}${var}, ${var}", "", ), @@ -908,7 +908,7 @@ class TestVarCommands(TestPromptPostProcessorBase): def test_cmd_set_ifundefined_if(self): # set, ifundefined and if commands self.process( - PromptPair( + InputTuple( "valuethis test is OKnot OK", "", ), @@ -917,7 +917,7 @@ class TestVarCommands(TestPromptPostProcessorBase): def test_cmd_set_ifundefined_if_2(self): # set, ifundefined and if commands self.process( - PromptPair( + InputTuple( "valuevalue2this test is OKnot OK", "", ), @@ -926,7 +926,7 @@ class TestVarCommands(TestPromptPostProcessorBase): def test_cmd_set_ifundefined_DP_if(self): # set, ifundefined (DP format) and if commands self.process( - PromptPair( + InputTuple( "${v?=value}this test is OKnot OK", "", ), @@ -935,7 +935,7 @@ class TestVarCommands(TestPromptPostProcessorBase): def test_cmd_set_ifundefined_DP_if_2(self): # set, ifundefined (DP format) and if commands self.process( - PromptPair( + InputTuple( "${v=!value}${v?=!value2}this test is OKnot OK", "", ), @@ -944,7 +944,7 @@ class TestVarCommands(TestPromptPostProcessorBase): def test_cmd_echo_sysvar(self): self.process( - PromptPair( + InputTuple( "${_model:defval}", "", ), @@ -953,7 +953,7 @@ class TestVarCommands(TestPromptPostProcessorBase): def test_cmd_ext(self): # ext self.process( - PromptPair( + InputTuple( "trigger1trigger2trigger4trigger5", "", ), @@ -965,7 +965,7 @@ class TestVarCommands(TestPromptPostProcessorBase): def test_cmd_ext_map_notrigger(self): # ext mapping, no trigger self.process( - PromptPair( + InputTuple( "", "", ), @@ -974,7 +974,7 @@ class TestVarCommands(TestPromptPostProcessorBase): def test_cmd_ext_map1(self): # ext mapping, no lora self.process( - PromptPair( + InputTuple( "inlinetrigger", "", ), @@ -983,7 +983,7 @@ class TestVarCommands(TestPromptPostProcessorBase): def test_cmd_ext_map2(self): # ext mapping, lora with weight self.process( - PromptPair( + InputTuple( "inlinetrigger", "", ), @@ -1004,7 +1004,7 @@ class TestVarCommands(TestPromptPostProcessorBase): def test_cmd_ext_map3(self): # ext mapping, lora with weight adjusted self.process( - PromptPair( + InputTuple( "inlinetrigger", "", ), @@ -1025,7 +1025,7 @@ class TestVarCommands(TestPromptPostProcessorBase): def test_cmd_ext_map4(self): # ext mapping, lora with parameters self.process( - PromptPair( + InputTuple( "inlinetrigger", "", ), @@ -1046,7 +1046,7 @@ class TestVarCommands(TestPromptPostProcessorBase): def test_cmd_ext_map5(self): # ext mapping, lora with no parameters self.process( - PromptPair( + InputTuple( "inlinetrigger", "", ), diff --git a/tests/tests_variants.py b/tests/tests_variants.py index ac06990..bca077d 100644 --- a/tests/tests_variants.py +++ b/tests/tests_variants.py @@ -2,7 +2,7 @@ from dataclasses import replace from ppp import PromptPostProcessor from ppp_classes import ONWARNING_CHOICES # type: ignore -from .base_tests import OutputTuple, PromptPair, TestPromptPostProcessorBase +from .base_tests import OutputTuple, InputTuple, TestPromptPostProcessorBase if __name__ == "__main__": raise SystemExit("This script must not be run directly") @@ -17,7 +17,7 @@ class TestModelVariants(TestPromptPostProcessorBase): def test_variants(self): self.process( - PromptPair( + InputTuple( "test1test2test3test4", "", ), diff --git a/tests/tests_wildcards.py b/tests/tests_wildcards.py index eac3264..4fb373c 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 OutputTuple, PromptPair, TestPromptPostProcessorBase +from .base_tests import OutputTuple, InputTuple, TestPromptPostProcessorBase if __name__ == "__main__": raise SystemExit("This script must not be run directly") @@ -17,7 +17,7 @@ class TestWildcards(TestPromptPostProcessorBase): def test_wc_ignore(self): # wildcards with ignore option self.process( - PromptPair("__bad_wildcard__", "{option1|option2}"), + InputTuple("__bad_wildcard__", "{option1|option2}"), OutputTuple("__bad_wildcard__", "{option1|option2}"), ppp=PromptPostProcessor( self.ppp_logger, @@ -36,7 +36,7 @@ class TestWildcards(TestPromptPostProcessorBase): def test_wc_remove(self): # wildcards with remove option self.process( - PromptPair( + InputTuple( "[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}", ), @@ -61,7 +61,7 @@ class TestWildcards(TestPromptPostProcessorBase): def test_wc_warn(self): # wildcards with warn option self.process( - PromptPair("__bad_wildcard__", "{option1|option2}"), + InputTuple("__bad_wildcard__", "{option1|option2}"), OutputTuple(PromptPostProcessor.WILDCARD_WARNING + "__bad_wildcard__", "{option1|option2}"), ppp=PromptPostProcessor( self.ppp_logger, @@ -80,7 +80,7 @@ class TestWildcards(TestPromptPostProcessorBase): def test_wc_stop(self): # wildcards with stop option self.process( - PromptPair("__bad_wildcard__", "{option1|option2}"), + InputTuple("__bad_wildcard__", "{option1|option2}"), OutputTuple( PromptPostProcessor.WILDCARD_STOP.format("__bad_wildcard__") + "__bad_wildcard__", "{option1|option2}", @@ -103,7 +103,7 @@ class TestWildcards(TestPromptPostProcessorBase): def test_wcinvar_warn(self): # wildcards in var with warn option self.process( - PromptPair("${v=__bad_wildcard__}${v}", ""), + InputTuple("${v=__bad_wildcard__}${v}", ""), OutputTuple(PromptPostProcessor.WILDCARD_WARNING + "__bad_wildcard__", ""), ppp=PromptPostProcessor( self.ppp_logger, @@ -122,161 +122,161 @@ class TestWildcards(TestPromptPostProcessorBase): def test_wc_invalid_name(self): self.process( - PromptPair("the choices are: ___invalid__", ""), + InputTuple("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__", ""), + InputTuple("the choices are: __text/wildcard1__", ""), OutputTuple("the choices are: choice2", ""), ppp="nocup", ) def test_wc_wildcard1a_json(self): # simple json wildcard self.process( - PromptPair("the choices are: __json/wildcard1__", ""), + InputTuple("the choices are: __json/wildcard1__", ""), OutputTuple("the choices are: choice2", ""), ppp="nocup", ) def test_wc_wildcard1a_yaml(self): # simple yaml wildcard self.process( - PromptPair("the choices are: __yaml/wildcard1__", ""), + InputTuple("the choices are: __yaml/wildcard1__", ""), 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__", ""), + InputTuple("the choices are: __2-$$text/wildcard1__", ""), 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__", ""), + InputTuple("the choices are: __2-$$json/wildcard1__", ""), 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__", ""), + InputTuple("the choices are: __2-$$yaml/wildcard1__", ""), 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__", ""), + InputTuple("the choices are: __text/wildcard2__", ""), 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__", ""), + InputTuple("the choices are: __json/wildcard2__", ""), 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__", ""), + InputTuple("the choices are: __yaml/wildcard2__", ""), 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__", ""), + InputTuple("the choice is: __testwc/test2__", ""), OutputTuple("the choice is: 2", ""), ppp="nocup", ) def test_wc_test3_yaml(self): # simple yaml wildcard self.process( - PromptPair("the choice is: __testwc/test3__", ""), + InputTuple("the choice is: __testwc/test3__", ""), 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'__", ""), + InputTuple("the choice is: __yaml/wildcard2'2'__", ""), 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'__", ""), + InputTuple("the choice is: __yaml/wildcard2'2-3'__", ""), 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'__", ""), + InputTuple("the choice is: __yaml/wildcard2'label1'__", ""), 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'__", ""), + InputTuple("the choice is: __yaml/wildcard2'label2'__", ""), 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'__", ""), + InputTuple("the choice is: __yaml/wildcard2'label1,label2'__", ""), 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'__", ""), + InputTuple("the choice is: __yaml/wildcard2'2,label2'__", ""), 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'__", ""), + InputTuple("the choice is: __yaml/wildcard2'label1+label3'__", ""), 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'__", ""), + InputTuple("the choice is: __yaml/wildcard2bis'#label1+label3'__", ""), 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'__", ""), + InputTuple("the choice is: __yaml/wildcard2bisbis'#label1+label3'__", ""), 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'__", ""), + InputTuple("${v=label1}the choice is: __yaml/wildcard2bisbis'#${v}+label3'__", ""), OutputTuple("the choice is: choice1bisbis", ""), ppp="nocup", ) def test_wc_wildcard_default_filter(self): # wildcard with default filter self.process( - PromptPair( + InputTuple( "the choice is: __yaml/wildcard2__, __yaml/wildcard2__", "", ), @@ -286,7 +286,7 @@ class TestWildcards(TestPromptPostProcessorBase): def test_wc_wildcard_default_filter2(self): # wildcard with default filter with variable self.process( - PromptPair( + InputTuple( "${v=label1}the choice is: __yaml/wildcard2__, __yaml/wildcard2__", "", ), @@ -296,49 +296,49 @@ class TestWildcards(TestPromptPostProcessorBase): def test_wc_nested_wildcard_text(self): # nested text wildcard with repeating multiple choices self.process( - PromptPair("the choices are: __r3$$-$$text/wildcard3__", ""), + InputTuple("the choices are: __r3$$-$$text/wildcard3__", ""), 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__", ""), + InputTuple("the choices are: __r3$$-$$json/wildcard3__", ""), 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__", ""), + InputTuple("the choices are: __r3$$-$$yaml/wildcard3__", ""), 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__", ""), + InputTuple("the choices are: __yaml/empty_wildcard__", ""), 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__", ""), + InputTuple("the choices are: __yaml/wildcard4__", ""), 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__", ""), + InputTuple("the choices are: __yaml/wildcard6__", ""), OutputTuple("the choices are: choice2", ""), ppp="nocup", ) def test_wc_choice_wildcard_mix(self): # choices with wildcard mix self.process( - PromptPair("the choices are: {__~2$$yaml/wildcard2__|choice0}", ""), + InputTuple("the choices are: {__~2$$yaml/wildcard2__|choice0}", ""), [ OutputTuple("the choices are: choice0", ""), OutputTuple("the choices are: choice1, choice3", ""), @@ -349,7 +349,7 @@ class TestWildcards(TestPromptPostProcessorBase): def test_wc_unsupportedsampler(self): # unsupported sampler self.process( - PromptPair("the choices are: __@yaml/wildcard2__", ""), + InputTuple("the choices are: __@yaml/wildcard2__", ""), OutputTuple("", ""), ppp="nocup", interrupted=True, @@ -357,42 +357,42 @@ class TestWildcards(TestPromptPostProcessorBase): def test_wc_wildcard_globbing(self): # wildcard with globbing self.process( - PromptPair("the choices are: __yaml/wildcard[12]__, __yaml/wildcard?__", ""), + InputTuple("the choices are: __yaml/wildcard[12]__, __yaml/wildcard?__", ""), 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__", ""), + InputTuple("the choices are: __yaml/wildcard5(var=test)__, __yaml/wildcard5__", ""), 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__", ""), + InputTuple("the choices are: __yaml/wildcardPS__", ""), 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__", ""), + InputTuple("the choices are: __yaml/anonwildcards__", ""), 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__", ""), + InputTuple("the choices are: __yaml_input/wildcardI__", ""), OutputTuple("the choices are: choice2", ""), ppp="nocup", ) def test_wc_circular(self): # wildcard circular reference self.process( - PromptPair("the choices are: __yaml/circular1__", ""), + InputTuple("the choices are: __yaml/circular1__", ""), OutputTuple("", ""), ppp="nocup", interrupted=True, @@ -400,14 +400,14 @@ class TestWildcards(TestPromptPostProcessorBase): def test_wc_including(self): # wildcard including another wildcard self.process( - PromptPair("the choices are: __yaml/including__", ""), + InputTuple("the choices are: __yaml/including__", ""), 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__", ""), + InputTuple("the choices are: __yaml/including1__", ""), OutputTuple("", ""), ppp="nocup", interrupted=True, @@ -415,7 +415,7 @@ class TestWildcards(TestPromptPostProcessorBase): def test_wc_dynamicwildcard(self): # wildcard built from variables self.process( - PromptPair( + InputTuple( "the choices are: ${x={1|2|3}}${w=yaml/wildcard${x}}__yaml/wildcard${x}__ __${w}__ ____", "", ), @@ -427,7 +427,7 @@ class TestWildcards(TestPromptPostProcessorBase): def test_wc_combinatorial_1(self): # combinatorial wildcard with variable self.process( - PromptPair("the choices are: __2$$yaml/wildcard2__, ${v:{option1|option2}}", ""), + InputTuple("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"}), @@ -447,7 +447,7 @@ class TestWildcards(TestPromptPostProcessorBase): def test_wc_combinatorial_2(self): # combinatorial wildcard self.process( - PromptPair("__yaml/wildcard2__", ""), + InputTuple("__yaml/wildcard2__", ""), [ # 36 combinations # groups of 3 ## same choice repeated 3 times @@ -501,7 +501,7 @@ class TestWildcards(TestPromptPostProcessorBase): def test_wc_combinatorial_3(self): # combinatorial wildcard (keep choice order) self.process( - PromptPair("__2-3$$-$$yaml/wildcard2__", ""), + InputTuple("__2-3$$-$$yaml/wildcard2__", ""), [ # 4 combinations # groups of 3 ## choices 1, 2, 3 @@ -532,7 +532,7 @@ class TestWildcards(TestPromptPostProcessorBase): def test_wc_combinatorial_4(self): # combinatorial wildcard (don't keep choice order) self.process( - PromptPair("__2-3$$-$$yaml/wildcard2__", ""), + InputTuple("__2-3$$-$$yaml/wildcard2__", ""), [ # 12 combinations # groups of 3 ## choices 1, 2, 3 in all positions @@ -559,7 +559,7 @@ class TestWildcards(TestPromptPostProcessorBase): def test_wc_combinatorial_5(self): # combinatorial nested wildcards and multiselection enmappings self.process( - PromptPair("{__yaml/wildcard1__|__yaml/wildcard3__|}", ""), + InputTuple("{__yaml/wildcard1__|__yaml/wildcard3__|}", ""), [ # 11 combinations # first wildcard OutputTuple("choice1", ""),