From fdd2f5ba51b9226661cd4dcb71334a7870bfce23 Mon Sep 17 00:00:00 2001 From: Antonio Cordero Balcazar Date: Thu, 9 Apr 2026 16:47:55 +0200 Subject: [PATCH] * Added result checks for special character sequences and unmatched parentheses/brackets. --- ppp.py | 46 ++++++++++++++++++++++++------ ppp_tree.py | 2 +- tests/tests_cleanup.py | 65 ++++++++++++++++++++++++++++++++++++++++++ 3 files changed, 103 insertions(+), 10 deletions(-) diff --git a/ppp.py b/ppp.py index 927f6be..82d5009 100644 --- a/ppp.py +++ b/ppp.py @@ -972,6 +972,43 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in 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 foundP = bool(p_processor.detectedWildcards) foundNP = bool(n_processor.detectedWildcards) @@ -994,15 +1031,6 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in self.WILDCARD_STOP.format(ppwl) if foundP else "", self.WILDCARD_STOP.format(npwl) if foundNP else "", ) - - # 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: - self.log( - logging.WARNING, - f"""Found probably invalid character sequences on the result ({', '.join(map(lambda x: '"' + x + '"', set(found_sequences)))}). Something might be wrong!""", - ) return prompt, negative_prompt, all_variables def process_prompt( diff --git a/ppp_tree.py b/ppp_tree.py index 673884b..f6f097d 100644 --- a/ppp_tree.py +++ b/ppp_tree.py @@ -1005,7 +1005,7 @@ class TreeProcessor(lark.visitors.Interpreter): wcs = self.__state.wildcards_obj.get_wildcards(cmd_args) if not wcs: self.warn_or_stop( - f"Not found included wildcard '{escape_single_quotes(cmd_args)}' at {msg_where}!" + f"Included wildcard '{escape_single_quotes(cmd_args)}' not found at {msg_where}!" ) c_weight = float(c.get("weight", 1.0)) for wc in wcs: diff --git a/tests/tests_cleanup.py b/tests/tests_cleanup.py index 8211771..cd8e644 100644 --- a/tests/tests_cleanup.py +++ b/tests/tests_cleanup.py @@ -1,3 +1,4 @@ +import logging from dataclasses import replace from ppp import PromptPostProcessor # pylint: disable=import-error @@ -132,3 +133,67 @@ class TestCleanup(TestPromptPostProcessorBase): ), ppp="nocup", ) + + # Result warnings tests + + 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", ""), + PromptPair("(unclosed paren", ""), + ) + self.assertTrue( + any("Unmatched" in msg for msg in cm.output), + "Expected an 'Unmatched' warning for open parenthesis", + ) + + 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)", ""), + PromptPair("extra close paren)", ""), + ) + self.assertTrue( + any("Unmatched" in msg for msg in cm.output), + "Expected an 'Unmatched' warning for close parenthesis", + ) + + 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]", ""), + PromptPair("(mismatched]", ""), + ) + self.assertTrue( + any("Mismatched" in msg or "Unmatched" in msg for msg in cm.output), + "Expected a 'Mismatched' or 'Unmatched' warning for bracket mismatch", + ) + + 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", ""), + PromptPair("unclosed [bracket", ""), + ) + self.assertTrue( + any("Unmatched" in msg for msg in cm.output), + "Expected an 'Unmatched' warning for open bracket", + ) + + 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))", ""), + PromptPair("[(unmatched [bracket))", ""), + ) + self.assertTrue( + any("Unmatched" in msg for msg in cm.output), + "Expected an 'Unmatched' warning", + ) + + 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\]", ""), + PromptPair(r"text with \(escaped unmatched\]", ""), + )