* Added result checks for special character sequences and unmatched parentheses/brackets.

This commit is contained in:
Antonio Cordero Balcazar
2026-04-09 16:47:55 +02:00
parent 62939f4416
commit fdd2f5ba51
3 changed files with 103 additions and 10 deletions
+37 -9
View File
@@ -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(
+1 -1
View File
@@ -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:
+65
View File
@@ -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\]", ""),
)