* Added result checks for special character sequences and unmatched parentheses/brackets.
This commit is contained in:
@@ -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
@@ -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:
|
||||
|
||||
@@ -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\]", ""),
|
||||
)
|
||||
|
||||
Reference in New Issue
Block a user