Files
acorderob-sd-webui-prompt-p…/ppp_tree.py
T

1631 lines
78 KiB
Python

from collections import namedtuple
import math
import re
import textwrap
import time
from typing import Optional
import lark
import numpy as np
from ppp_classes import IFWILDCARDS_CHOICES, PPPState
from ppp_enmappings import PPPENMappingVariant
from ppp_logging import DEBUG_LEVEL
from ppp_utils import escape_single_quotes, format_output
from ppp_common import parse_prompt, warn_or_stop
from ppp_wildcards import PPPWildcard
class TreeProcessor(lark.visitors.Interpreter):
"""
A class for interpreting and processing a tree generated by the prompt parser.
Args:
state (PPPState): The state object containing the current processing state.
rng (numpy.random.Generator): The random number generator.
Attributes:
add_at (dict): The dictionary to store the content to be added at different positions of the negative prompt.
insertion_at (list): The list of insertion points in the negative prompt.
detectedWildcards (list): The list of detected invalid wildcards or choices.
result (str): The final processed prompt.
"""
def __init__(self, state: PPPState, rng: np.random.Generator):
super().__init__()
self.__state = state
self.__debug_level = state.options.debug_level
self.__rng = rng
self.AccumulatedShell = namedtuple("AccumulatedShell", ["type", "data"])
self.NegTag = namedtuple("NegTag", ["start", "end", "content", "parameters", "shell"])
self.__shell: list[self.AccumulatedShell] = [] # type: ignore
self.__negtags: list[self.NegTag] = [] # type: ignore
self.__already_processed: list[str] = []
self.__is_negative = False
self.__wildcard_filters = {}
self.__seen_wildcards: list[str] = []
self.add_at: dict = {"start": [], "insertion_point": [[] for x in range(10)], "end": []}
self.insertion_at: list[tuple[int, int]] = [None for x in range(10)]
self.detectedWildcards: list[str] = []
self.result = ""
def warn_or_stop(self, message: str, e: Exception = None):
warn_or_stop(self.__state, self.__is_negative, message, e)
def start_visit(self, prompt_description: str, parsed_prompt: lark.Tree, is_negative: bool = False) -> str:
"""
Start the visit process.
Args:
prompt_description (str): The description of the prompt.
parsed_prompt (Tree): The parsed prompt.
is_negative (bool): Whether the prompt is negative or not.
Returns:
str: The processed prompt.
"""
t1 = time.monotonic_ns()
self.__is_negative = is_negative
if self.__debug_level != DEBUG_LEVEL.none:
self.__state.logger.info(f"Processing {prompt_description}...")
self.visit(parsed_prompt)
t2 = time.monotonic_ns()
if self.__debug_level != DEBUG_LEVEL.none:
self.__state.logger.info(f"Process {prompt_description} time: {(t2 - t1) / 1_000_000_000:.3f} seconds")
return self.result
def __visit(
self,
node: lark.Tree | lark.Token | list[lark.Tree | lark.Token] | None,
restore_state: bool = False,
discard_content: bool = False,
) -> str:
"""
Visit a node in the tree and process it or accumulate its value if it is a Token.
Args:
node (Tree|Token|list): The node or list of nodes to visit.
restore_state (bool): Whether to restore the state after visiting the node.
discard_content (bool): Whether to discard the content of the node.
Returns:
str: The result of the visit.
"""
backup_result = self.result
# if self.__debug_level == DEBUG_LEVEL.full:
# self.__state.logger.debug(f"Visiting node {node}.")
if restore_state:
# if self.__debug_level == DEBUG_LEVEL.full:
# self.__state.logger.debug("Backing up state before visiting.")
backup_shell = self.__shell.copy()
backup_negtags = self.__negtags.copy()
backup_already_processed = self.__already_processed.copy()
backup_add_at = self.add_at.copy()
backup_insertion_at = self.insertion_at.copy()
backup_detectedwildcards = self.detectedWildcards.copy()
backup_user_variables = self.__state.user_variables.copy()
backup_echoed_variables = self.__state.echoed_variables.copy()
if node is not None:
if isinstance(node, list):
for child in node:
self.__visit(child)
elif isinstance(node, lark.Tree):
self.visit(node)
elif isinstance(node, lark.Token):
self.result += node
len_backup = len(backup_result)
# if self.result[:len_backup] == backup_result: # this is only necessary if we call parse_prompt with a parser from "start", because it resets the result
added_result = self.result[len_backup:]
# else:
# added_result = self.result
if discard_content or restore_state:
self.result = backup_result
if restore_state:
# if self.__debug_level == DEBUG_LEVEL.full:
# self.__state.logger.debug("Restoring state after visiting.")
self.__shell = backup_shell
self.__negtags = backup_negtags
self.__already_processed = backup_already_processed
self.add_at = backup_add_at
self.insertion_at = backup_insertion_at
self.detectedWildcards = backup_detectedwildcards
self.__state.user_variables.clear()
self.__state.user_variables.update(backup_user_variables)
self.__state.echoed_variables.clear()
self.__state.echoed_variables.update(backup_echoed_variables)
return added_result
def __get_original_node_content(self, node: lark.Tree | lark.Token, default=None) -> str:
"""
Get the original content of a node.
Args:
node (Tree|Token): The node to get the content from.
default: The default value to return if the content is not found.
Returns:
str: The original content of the node.
"""
return node.meta.content if hasattr(node, "meta") and node.meta is not None and not node.meta.empty else default
def __get_user_variable_value(self, name: str, evaluate=True, visit=False) -> str:
"""
Get the value of a user variable.
Args:
name (str): The name of the user variable.
evaluate (bool): Whether to evaluate the variable.
visit (bool): Whether to also visit the variable (add to result).
Returns:
str: The value of the user variable.
"""
v = self.__state.user_variables.get(name, None)
if v is not None:
visited = False
if isinstance(v, lark.Tree):
if evaluate:
v = self.__visit(v, not visit)
visited = visit
else:
v = self.__get_original_node_content(v, "not evaluated yet")
if visit and not visited:
self.result += v
return v
def get_final_user_variable(self, name: str) -> str:
return self.__get_user_variable_value(name, True, False)
def __set_user_variable_value(self, name: str, value: str):
"""
Set the value of a user variable.
Args:
name (str): The name of the user variable.
value (str): The value to be set.
"""
self.__state.user_variables[name] = value
def __remove_user_variable(self, name: str):
"""
Remove a user variable.
Args:
name (str): The name of the user variable.
"""
if name in self.__state.user_variables:
del self.__state.user_variables[name]
def __debug_end(self, construct: str, start_result: str, duration: int, info=None):
"""
Log the end of a construct processing.
Args:
construct (str): The name of the construct.
start_result (str): The initial result.
duration (int): The duration of the processing in ns.
info: Additional information to log.
"""
if self.__debug_level == DEBUG_LEVEL.full:
info = f"({info}) " if info is not None and info != "" else ""
output = self.result[len(start_result) :]
if output != "":
output = f" >> '{escape_single_quotes(output)}'"
self.__state.logger.debug(
format_output(f"TreeProcessor.{construct} {info}({duration / 1_000_000_000:.3f} seconds){output}")
)
def __resolve_cond_value(self, c: str):
"""Resolve a condition value: try int first, fall back to variable lookup."""
try:
return int(c)
except ValueError:
# Bare identifier - resolve as variable reference
if c.startswith("_"):
val = self.__state.system_variables.get(c, None)
if val is None:
val = ""
self.warn_or_stop(f"Unknown system variable {c}")
else:
val = self.__get_user_variable_value(c)
if val is None:
val = ""
self.warn_or_stop(f"Unknown user variable {c}")
return val.lower() if isinstance(val, str) else val
def __eval_basiccondition(self, cond_var: str, cond_comp: str, cond_value: str | list[str]) -> bool:
"""
Evaluate a condition based on the given variable, comparison, and value.
Args:
cond_var (str): The variable to be compared.
cond_comp (str): The comparison operator.
cond_value (str or list[str]): The value to be compared with.
Returns:
bool: The result of the condition evaluation.
"""
if cond_var.lower() == "false":
var_value = "false"
elif cond_var.lower() == "true":
var_value = "true"
elif cond_var.startswith("_"): # system variable
var_value = self.__state.system_variables.get(cond_var, None)
if var_value is None:
var_value = ""
self.warn_or_stop(f"Unknown system variable {cond_var}")
else: # user variable
var_value = self.__get_user_variable_value(cond_var)
if var_value is None:
var_value = ""
self.warn_or_stop(f"Unknown user variable {cond_var}")
if isinstance(var_value, str):
var_value = var_value.lower()
if isinstance(cond_value, list):
comp_ops = {
"contains": lambda x, y: y in x,
"in": lambda x, y: x == y,
}
else:
cond_value = [cond_value]
comp_ops = {
"eq": lambda x, y: x == y,
"ne": lambda x, y: x != y,
"gt": lambda x, y: x > y,
"lt": lambda x, y: x < y,
"ge": lambda x, y: x >= y,
"le": lambda x, y: x <= y,
"contains": lambda x, y: y in x,
"truthy": lambda x, y: bool(x),
}
if cond_comp not in comp_ops:
return False
cond_value_adjusted = list(
(
c[1:-1].lower()
if c.startswith('"') or c.startswith("'")
else (
True
if c.lower() == "true"
else False if c.lower() == "false" or c == "" else self.__resolve_cond_value(c)
)
)
for c in cond_value
)
result = False
for c in cond_value_adjusted:
if isinstance(c, str):
var_value_adjusted = var_value
elif isinstance(c, bool) and var_value != "false" and var_value != "" and var_value is not False:
var_value_adjusted = True
elif isinstance(c, bool) and (var_value != "true" or var_value is False):
var_value_adjusted = False
else:
try:
var_value_adjusted = int(var_value)
except (ValueError, TypeError):
self.warn_or_stop(
f"Cannot convert variable value '{escape_single_quotes(var_value)}' to integer for comparison"
)
return False
result = comp_ops[cond_comp](var_value_adjusted, c)
if result:
break
return result
def __eval_condition(self, condition: lark.Tree) -> bool:
"""
Evaluate an if condition based on the given condition tree.
Args:
condition (Node): The condition tree to be evaluated.
Returns:
bool: The result of the if condition evaluation.
"""
# self.__state.logger.debug(f"__eval_condition {condition.data}")
if condition.data == "operation_and":
cond_result = True
for c in condition.children:
cond_result = cond_result and self.__eval_condition(c)
if not cond_result:
break
elif condition.data == "operation_or":
cond_result = False
for c in condition.children:
cond_result = cond_result or self.__eval_condition(c)
if cond_result:
break
elif condition.data == "operation_not":
cond_result = not self.__eval_condition(condition.children[0])
else: # truthy_operand / comparison_simple_value / comparison_list_value
# we get the name of the variable
cond_var = str(condition.children[0])
poscomp = 1
invert = False
if poscomp >= len(condition.children):
# no condition, just a variable
cond_comp = "truthy"
cond_value = "true"
else:
# we get the comparison (with possible not) and the value
cond_comp = str(condition.children[poscomp])
if cond_comp == "not":
invert = not invert
poscomp += 1
cond_comp = str(condition.children[poscomp])
poscomp += 1
cond_value_node = condition.children[poscomp]
cond_value = (
list(str(v) for v in cond_value_node.children)
if isinstance(cond_value_node, (lark.Tree, list))
else str(cond_value_node) if isinstance(cond_value_node, lark.Token) else cond_value_node
)
cond_result = self.__eval_basiccondition(cond_var, cond_comp, cond_value)
if invert:
cond_result = not cond_result
return cond_result
def promptcomp(self, tree: lark.Tree):
"""
Process a prompt composition construct in the tree.
"""
start_result = self.result
t1 = time.monotonic_ns()
self.__visit(tree.children[0])
and_processing = self.__state.host_config.get("and", "ok")
if len(tree.children) > 1:
and_replacements = {
"eol": ("replaced with EOL", "\n"),
"comma": ("replaced with COMMA", ", "),
"remove": ("removed", " "),
}
if tree.children[1] is not None:
self.result += f":{tree.children[1]}"
for i in range(2, len(tree.children), 3):
if and_processing in and_replacements.keys():
self.result = (
self.result.rstrip()
+ and_replacements[and_processing][1]
+ self.__visit(tree.children[i + 1], False, True).lstrip()
)
if self.__debug_level == DEBUG_LEVEL.full:
self.__state.logger.debug(f"AND construct {and_replacements[and_processing][0]}")
elif and_processing == "error":
self.warn_or_stop("AND constructs are not allowed!")
else: # and_processing == "ok":
if self.__state.options.cup_ands:
self.result = re.sub(r"[, ]+$", "\n" if self.__state.options.cup_ands_eol else " ", self.result)
if self.result[-1:].isalnum(): # add space if needed
self.result += " "
self.result += "AND"
added_result = self.__visit(tree.children[i + 1], False, True)
if self.__state.options.cup_ands:
added_result = re.sub(r"^[, ]+", " ", added_result)
if added_result[0:1].isalnum(): # add space if needed
added_result = " " + added_result
self.result += added_result
if tree.children[i + 2] is not None:
self.result += f":{tree.children[i+2]}"
t2 = time.monotonic_ns()
self.__debug_end("promptcomp", start_result, t2 - t1)
def scheduled(self, tree: lark.Tree):
"""
Process a scheduling construct in the tree and add it to the accumulated shell.
"""
start_result = self.result
t1 = time.monotonic_ns()
before = tree.children[0]
after = tree.children[-2]
pos_str = tree.children[-1]
pos = float(pos_str)
if pos >= 1:
pos = int(pos)
scheduling_processing = self.__state.host_config.get("scheduling", "ok")
if scheduling_processing == "before":
if self.__debug_level == DEBUG_LEVEL.full:
self.__state.logger.debug("Scheduling construct removed, taking before option")
if before is not None:
self.__visit(before)
elif scheduling_processing == "after":
if self.__debug_level == DEBUG_LEVEL.full:
self.__state.logger.debug("Scheduling construct removed, taking after option")
if after is not None:
self.__visit(after)
elif scheduling_processing == "first":
if self.__debug_level == DEBUG_LEVEL.full:
self.__state.logger.debug("Scheduling construct removed, taking first option")
if before is not None:
self.__visit(before)
elif after is not None:
self.__visit(after)
elif scheduling_processing == "remove":
if self.__debug_level == DEBUG_LEVEL.full:
self.__state.logger.debug("Scheduling construct removed")
elif scheduling_processing == "error":
self.warn_or_stop("Scheduling constructs are not allowed!")
else: # scheduling_processing == "ok"
# self.__shell.append(self.AccumulatedShell("sc", pos))
self.result += "["
if before is not None:
if self.__debug_level == DEBUG_LEVEL.full:
self.__state.logger.debug(f"Shell scheduled before with position {pos}")
self.__shell.append(self.AccumulatedShell("scb", pos))
self.__visit(before)
self.__shell.pop()
if self.__debug_level == DEBUG_LEVEL.full:
self.__state.logger.debug(f"Shell scheduled after with position {pos}")
self.__shell.append(self.AccumulatedShell("sca", pos))
self.result += ":"
self.__visit(after)
self.__shell.pop()
if self.__state.options.cup_emptyconstructs and re.fullmatch(
re.escape(start_result) + r"\[:\s*", self.result
):
self.result = start_result
else:
self.result += f":{pos_str}]"
# self.__shell.pop()
t2 = time.monotonic_ns()
self.__debug_end("scheduled", start_result, t2 - t1, pos_str)
def alternate(self, tree: lark.Tree):
"""
Process an alternation construct in the tree and add it to the accumulated shell.
"""
start_result = self.result
t1 = time.monotonic_ns()
alternation_processing = self.__state.host_config.get("alternation", "ok")
if alternation_processing == "first":
if self.__debug_level == DEBUG_LEVEL.full:
self.__state.logger.debug("Alternation construct removed, taking first option")
self.__visit(tree.children[0])
elif alternation_processing == "remove":
if self.__debug_level == DEBUG_LEVEL.full:
self.__state.logger.debug("Alternation construct removed")
elif alternation_processing == "error":
self.warn_or_stop("Alternation constructs are not allowed!")
else: # alternation_processing == "ok"
# self.__shell.append(self.AccumulatedShell("al", len(tree.children)))
self.result += "["
for i, opt in enumerate(tree.children):
if self.__debug_level == DEBUG_LEVEL.full:
self.__state.logger.debug(f"Shell alternate option {i+1}")
self.__shell.append(self.AccumulatedShell("alo", {"pos": i + 1, "len": len(tree.children)}))
if i > 0:
self.result += "|"
self.__visit(opt)
self.__shell.pop()
self.result += "]"
if self.__state.options.cup_emptyconstructs and re.fullmatch(
re.escape(start_result) + r"\[\s*\]", self.result
):
self.result = start_result
# self.__shell.pop()
t2 = time.monotonic_ns()
self.__debug_end("alternate", start_result, t2 - t1)
def attention(self, tree: lark.Tree):
"""
Process a attention change construct in the tree and add it to the accumulated shell.
"""
start_result = self.result
t1 = time.monotonic_ns()
# weight_kind: -1: remove, 0=none, 1=decrease, 2=increase, 3=specific
if len(tree.children) == 2:
weight_str = tree.children[-1]
if weight_str is not None:
weight_kind = 3 # specific weight
weight = float(weight_str)
else:
weight_kind = 2 # increase attention
weight = 1.1
weight_str = "1.1"
else:
weight_kind = 1 # decrease attention
weight = 0.9
weight_str = "0.9"
if self.__debug_level == DEBUG_LEVEL.full:
self.__state.logger.debug(f"Shell attention with weight {weight}")
current_tree = tree.children[0]
if self.__state.options.cup_mergeattention:
while isinstance(current_tree, lark.Tree) and current_tree.data == "attention":
# we merge the weights
if len(current_tree.children) == 2:
inner_weight = current_tree.children[-1]
if inner_weight is not None:
inner_weight = float(inner_weight)
else:
inner_weight = 1.1
else:
inner_weight = 0.9
weight *= inner_weight
current_tree = current_tree.children[0]
weight = math.floor(weight * 100) / 100 # we round to 2 decimals
weight_str = f"{weight:.2f}".rstrip("0").rstrip(".")
if weight_str == "0.9":
weight_kind = 1
elif weight_str == "1.1":
weight_kind = 2
else:
weight_kind = 3
attention_processing = self.__state.host_config.get("attention", "ok")
if attention_processing == "parentheses":
if weight_kind == 1:
weight_kind = 3
weight_str = "0.9"
if self.__debug_level == DEBUG_LEVEL.full:
self.__state.logger.debug("Converted to parentheses format")
elif attention_processing == "disable":
weight_kind = 0
if self.__debug_level == DEBUG_LEVEL.full:
self.__state.logger.debug("Attention construct disabled")
elif attention_processing == "remove":
weight_kind = -1
if self.__debug_level == DEBUG_LEVEL.full:
self.__state.logger.debug("Attention construct removed")
elif attention_processing == "error":
self.warn_or_stop("Attention constructs are not allowed!")
# else: attention_processing == "ok":
if weight_kind == -1:
# we just ignore the attention construct
pass
elif weight_kind == 0:
# we just visit the content without adding any attention
self.__visit(current_tree)
else:
self.__shell.append(self.AccumulatedShell("at", (weight_kind, weight_str)))
if weight_kind == 1:
starttag = "["
self.result += starttag
self.__visit(current_tree)
endtag = "]"
elif weight_kind == 2:
starttag = "("
self.result += starttag
self.__visit(current_tree)
endtag = ")"
else: # weight_kind == 3
starttag = "("
self.result += starttag
self.__visit(current_tree)
endtag = f":{weight_str})"
if self.__state.options.cup_emptyconstructs and re.fullmatch(
re.escape(start_result + starttag) + r"\s*", self.result
):
self.result = start_result
else:
self.result += endtag
self.__shell.pop()
t2 = time.monotonic_ns()
self.__debug_end("attention", start_result, t2 - t1, weight_str)
def commandstn(self, tree: lark.Tree):
"""
Process a send to negative command in the tree and add it to the list of negative tags.
"""
start_result = self.result
info = None
t1 = time.monotonic_ns()
if not self.__is_negative:
negtagparameters = tree.children[0]
if negtagparameters is not None:
parameters = str(negtagparameters)
else:
parameters = ""
content = self.__visit(tree.children[1::], False, True)
self.__negtags.append(
self.NegTag(len(self.result), len(self.result), content, parameters, self.__shell.copy())
)
info = f"with {escape_single_quotes(parameters) or 'no parameters'} : {escape_single_quotes(content)}"
else:
self.warn_or_stop("Ignored negative command in negative prompt")
self.__visit(tree.children[1::])
t2 = time.monotonic_ns()
self.__debug_end("commandstn", start_result, t2 - t1, info)
def commandstni(self, tree: lark.Tree):
"""
Process a send to negative insertion point command in the tree and add it to the list of negative tags.
"""
start_result = self.result
info = None
t1 = time.monotonic_ns()
if self.__is_negative:
negtagparameters = tree.children[0]
if negtagparameters is not None:
parameters = str(negtagparameters)
else:
parameters = ""
self.__negtags.append(self.NegTag(len(self.result), len(self.result), "", parameters, self.__shell.copy()))
info = f"with {parameters or 'no parameters'}"
else:
self.warn_or_stop("Ignored negative insertion point command in positive prompt")
t2 = time.monotonic_ns()
self.__debug_end("commandstni", start_result, t2 - t1, info)
def __varset(
self,
command: str,
variable: str,
modifiers: lark.Tree | None,
content: lark.Tree | None,
):
"""
Process a generic set command in the tree.
"""
t1 = time.monotonic_ns()
start_result = self.result
if variable.startswith("_"):
self.warn_or_stop(
f"Invalid variable name '{escape_single_quotes(variable)}' detected! System variables cannot be set."
)
return
info = variable
value_description = self.__get_original_node_content(content, None)
value = content
modifiers_str: list[str] = [str(m) for m in modifiers.children] if modifiers is not None else []
if any(item in modifiers_str for item in ["+", "add"]):
info += f" += '{escape_single_quotes(value_description or '')}'"
raw_oldvalue = self.__state.user_variables.get(variable, None)
if raw_oldvalue is None:
newvalue = value
self.warn_or_stop(f"Unknown variable {escape_single_quotes(variable)}")
elif isinstance(raw_oldvalue, str):
newvalue = lark.Tree(
lark.Token("RULE", "varvalue"),
[lark.Token("plain", raw_oldvalue), value],
# Meta should be {"content": raw_oldvalue + value},
)
else:
newvalue = lark.Tree(
lark.Token("RULE", "varvalue"),
[raw_oldvalue, value],
# Meta should be {"content": raw_oldvalue.meta.content + value.meta.content},
)
elif any(item in modifiers_str for item in ["?", "ifundefined"]):
info += f" ?= '{escape_single_quotes(value_description or '')}'"
raw_oldvalue = self.__state.user_variables.get(variable, None)
if raw_oldvalue is None:
newvalue = value
else:
info += " (not set)"
newvalue = None
else:
newvalue = value
if newvalue is not None:
if any(item in modifiers_str for item in ["!", "evaluate"]):
newvalue = self.__visit(newvalue, False, True)
info += " =! "
else:
info += " = "
self.__set_user_variable_value(variable, newvalue)
currentvalue = self.__get_user_variable_value(variable, False)
if currentvalue is None:
info += "not evaluated yet"
else:
info += f"'{escape_single_quotes(currentvalue)}'"
t2 = time.monotonic_ns()
self.__debug_end(command, start_result, t2 - t1, info)
def variableset(self, tree: lark.Tree):
"""
Process a DP set variable command in the tree and add it to the dictionary of variables.
"""
modifiers = tree.children[1] or lark.Tree(lark.Token("RULE", "variablesetmodifiers"), [])
immediate = tree.children[2]
if immediate is not None:
modifiers.children = modifiers.children.copy()
modifiers.children.append(immediate)
self.__varset("variableset", str(tree.children[0]), modifiers, tree.children[3])
def commandset(self, tree: lark.Tree):
"""
Process a set command in the tree and add it to the dictionary of variables.
"""
self.__varset("commandset", str(tree.children[0]), tree.children[1], tree.children[2])
def __varecho(self, command: str, variable: str, default: lark.Tree | None):
"""
Process a generic echo command in the tree.
"""
t1 = time.monotonic_ns()
start_result = self.result
default_value = None
# if default is not None:
# default_value = self.__visit(default, True) # for log
value = self.__get_user_variable_value(variable, True, True)
if value is None:
if default is not None:
if self.__debug_level == DEBUG_LEVEL.full:
self.__state.logger.debug(
f"Variable '{escape_single_quotes(variable)}' not found, using default value"
)
v = self.__visit(default, False, True)
self.result += v
default_value = v
self.__state.echoed_variables[variable] = v
else:
self.warn_or_stop(f"Unknown variable {escape_single_quotes(variable)}")
default_value = ""
self.__state.echoed_variables[variable] = ""
else:
self.__state.echoed_variables[variable] = value
t2 = time.monotonic_ns()
info = variable
if default_value is not None:
info += f" with default '{escape_single_quotes(default_value)}'"
self.__debug_end(command, start_result, t2 - t1, info)
def variableuse(self, tree: lark.Tree):
"""
Process a DP use variable command in the tree.
"""
self.__varecho("variableuse", str(tree.children[0]), tree.children[1] if len(tree.children) > 1 else None)
def commandecho(self, tree: lark.Tree):
"""
Process an echo command in the tree.
"""
self.__varecho("commandecho", str(tree.children[0]), tree.children[1] if len(tree.children) > 1 else None)
def commandif(self, tree: lark.Tree):
"""
Process an if command in the tree.
"""
t1 = time.monotonic_ns()
start_result = self.result
for i, n in enumerate(tree.children):
content = n.children[-1]
if len(n.children) == 2: # its not an else
# has a condition
condition = n.children[0]
c = self.__get_original_node_content(condition, f"condition {i}")
if self.__eval_condition(condition):
self.__visit(content)
t2 = time.monotonic_ns()
self.__debug_end("commandif", start_result, t2 - t1, c)
return
else: # its an else
self.__visit(content)
t2 = time.monotonic_ns()
self.__debug_end("commandif", start_result, t2 - t1, "else")
return
def commandext(self, tree: lark.Tree):
"""
Process an extranetwork command in the tree.
"""
t1 = time.monotonic_ns()
start_result = self.result
extnet = "(ignored)"
if not self.__state.options.rem_removeextranetworktags:
extnet_type: str = (tree.children[0].children[0] or "") + str(tree.children[0].children[1])
is_mapping = extnet_type.startswith("$")
if is_mapping:
extnet_type = extnet_type[1:]
extnet_id: str = str(tree.children[1])
if extnet_id.startswith("'") or extnet_id.startswith('"'):
extnet_id = extnet_id[1:-1]
extnet_id = re.sub(r"\\(.)", r"\1", extnet_id) # so we can escape some special characters
parameters: str = ""
parameters_defaulted = False
if tree.children[2]:
parameters = str(tree.children[2])
elif extnet_type in ("lora", "hypernet"):
parameters = "1"
parameters_defaulted = True
if parameters.startswith("'") or parameters.startswith('"'):
parameters = parameters[1:-1]
parameters_is_number = bool(re.match(r"^[-+]?\d*\.?\d+$", parameters or ""))
condition = tree.children[3]
if not condition or self.__eval_condition(condition):
extnet_id = f"{extnet_type}:{extnet_id}"
triggers = tree.children[4] if len(tree.children) > 4 else None
extra_triggers = None
compiled_extra_triggers = None
if is_mapping:
found = self.__state.extranetwork_mappings_obj.cached_mappings.get(extnet_id, None)
# we assume the conditions do not change inside the prompt
found_in_cache = found is not None
if found is None:
found_mappings: list[PPPENMappingVariant] = []
else_mapping = None
if self.__state.extranetwork_mappings_obj:
enmapping = self.__state.extranetwork_mappings_obj.extranetwork_mappings.get(
extnet_id, None
)
if enmapping:
for v in enmapping.variants:
if v.condition:
try:
cnd = parse_prompt(
self.__state,
"condition",
v.condition,
self.__state.parsers["condition"],
True,
)
except lark.exceptions.UnexpectedInput as e:
self.warn_or_stop(
f"Error parsing condition '{escape_single_quotes(v.condition)}' 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 v.condition:
found_mappings.append(v)
else:
else_mapping = v
if found_mappings:
found = found_mappings[
self.__rng.choice(
len(found_mappings),
p=[v.weight or 1 for v in found_mappings],
)
]
else:
found = else_mapping
self.__state.extranetwork_mappings_obj.cached_mappings[extnet_id] = found
if found:
if found.name:
if not found_in_cache and self.__debug_level != DEBUG_LEVEL.none:
self.__state.logger.info(
f"Mapping extranetwork '{escape_single_quotes(extnet_id)}' to '{escape_single_quotes(extnet_type)}:{escape_single_quotes(found.name)}'"
)
extnet_id = f"{extnet_type}:{found.name}"
f_parameters = found.parameters
if not f_parameters and extnet_type in ("lora", "hypernet"):
f_parameters = "1"
found_parameters_is_number = True
else:
found_parameters_is_number = f_parameters and bool(
re.match(r"^[-+]?\d*\.?\d+$", str(f_parameters) or "")
)
if found_parameters_is_number and parameters_is_number:
parameters = f"{float(f_parameters) * float(parameters):.2f}".rstrip("0").rstrip(".")
elif f_parameters is not None and parameters_defaulted:
parameters = f_parameters
elif found.triggers:
if not found_in_cache and self.__debug_level != DEBUG_LEVEL.none:
self.__state.logger.info(
f"Mapping extranetwork '{escape_single_quotes(extnet_id)}' to just triggers"
)
extnet_id = None
else:
if not found_in_cache and self.__debug_level != DEBUG_LEVEL.none:
self.__state.logger.info(
f"Mapping extranetwork '{escape_single_quotes(extnet_id)}' to nothing"
)
extnet_id = None
if found.triggers:
extra_triggers = ", ".join(found.triggers)
try:
compiled_extra_triggers = parse_prompt(
self.__state, "triggers", extra_triggers, self.__state.parsers["content"], True
)
except lark.exceptions.UnexpectedInput as e:
self.warn_or_stop(
f"Error parsing triggers '{escape_single_quotes(extra_triggers)}' in extranetwork mapping '{escape_single_quotes(extnet_id)}'! : {e.__class__.__name__}",
e,
)
compiled_extra_triggers = None
else:
self.warn_or_stop(f"Extranetwork mapping '{escape_single_quotes(extnet_id)}' not found!")
if extnet_id:
extnet = f"<{extnet_id}:{parameters}>"
self.result += extnet
elif triggers or compiled_extra_triggers:
extnet = "(only triggers)"
if triggers or compiled_extra_triggers:
if extnet_id:
if not self.__state.options.cup_extranetworktags:
self.result += " "
else:
self.result += ", "
if triggers:
self.result += self.__visit(triggers, True, True)
if compiled_extra_triggers:
if triggers:
self.result += ", "
self.result += self.__visit(compiled_extra_triggers, True, True)
if triggers or compiled_extra_triggers:
self.result += ", "
t2 = time.monotonic_ns()
self.__debug_end("commandext", start_result, t2 - t1, extnet)
def commandsetwcdeffilter(self, tree: lark.Tree):
"""
Process a setwcdeffilter (Set Wildcard Default Filter) command in the tree.
"""
t1 = time.monotonic_ns()
start_result = self.result
wildcard_key: str = self.__visit(tree.children[0].children[1], False, True)
selected_wildcards = [x.key for x in self.__state.wildcards_obj.get_wildcards(wildcard_key)]
if not selected_wildcards:
self.warn_or_stop(f"Wildcard '{escape_single_quotes(wildcard_key)}' not found for default filter setting!")
else:
filter_object = tree.children[1].children[1] if tree.children[1] is not None else None
if filter_object is None:
for wc in selected_wildcards:
if self.__debug_level == DEBUG_LEVEL.full:
self.__state.logger.debug(f"Removed default filter for wildcard '{escape_single_quotes(wc)}'")
self.__state.wildcards_obj.set_wildcard_default_filter(wc, None)
else:
filter_specifier = self.__extract_filter_specifiers(filter_object)
for wc in selected_wildcards:
if self.__debug_level == DEBUG_LEVEL.full:
self.__state.logger.debug(f"Set default filter for wildcard '{escape_single_quotes(wc)}'")
self.__state.wildcards_obj.set_wildcard_default_filter(wc, filter_specifier)
t2 = time.monotonic_ns()
self.__debug_end("commandsetwcdeffilter", start_result, t2 - t1)
def extranetworktag(self, tree: lark.Tree):
"""
Process an extra network construct in the tree.
"""
t1 = time.monotonic_ns()
start_result = self.result
if not self.__state.options.rem_removeextranetworktags:
self.result += f"<{tree.children[0]}"
self.__visit(tree.children[1])
self.result += ">"
t2 = time.monotonic_ns()
self.__debug_end("extranetworktag", start_result, t2 - t1)
def __get_choices_internal_get(
self,
choice_values: list[dict],
filter_specifier: Optional[list[list[str]]] = None,
wildcard_key: str = None,
) -> list[dict]:
msg_where = f"wildcard '{escape_single_quotes(wildcard_key)}'" if wildcard_key else "choices"
if filter_specifier is not None:
filtered_choice_values = []
for i, c in enumerate(choice_values):
passes = False
for or_ in filter_specifier:
tmp_pass = True
for and_ in or_:
if and_.isdecimal():
if int(and_) != i:
tmp_pass = False
break
elif re.match(r"^\d+-\d+$", and_):
try:
start, end = map(int, and_.split("-", 1))
if not start <= i <= end:
tmp_pass = False
break
except ValueError:
tmp_pass = False
break
elif and_.lower() not in c.get("labels", []):
tmp_pass = False
break
if tmp_pass:
passes = True
break
if passes:
filtered_choice_values.append(c)
if not filtered_choice_values:
self.warn_or_stop(
f"Wildcard filter specifier '{escape_single_quotes(','.join(['+'.join(y for y in x) for x in filter_specifier]))}' found no matches in choices for wildcard '{escape_single_quotes(wildcard_key)}'!"
)
else:
filtered_choice_values = choice_values.copy()
expanded_choice_values = []
for i, c in enumerate(filtered_choice_values):
if c.get("command", False):
content_text = self.__visit(c.get("content", ""), False, True).strip()
(cmd, cmd_args) = content_text.split()
if cmd == "include":
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}!"
)
c_weight = float(c.get("weight", 1.0))
for wc in wcs:
if wc.key in self.__seen_wildcards:
self.warn_or_stop(
f"Circular reference detected including wildcard '{escape_single_quotes(wc.key)}' at {msg_where} (chain starts at '{escape_single_quotes(self.__seen_wildcards[0])}')!"
)
continue
self.__seen_wildcards.append(wc.key)
if self.__debug_level == DEBUG_LEVEL.full:
self.__state.logger.debug(f"Seen wildcard '{escape_single_quotes(wc.key)}'")
self.__state.logger.debug(
f"Including choices from wildcard '{escape_single_quotes(wc.key)}'"
)
(_, choice_values) = self.__check_wildcard_initialization(wc)
if choice_values is not None:
ch_values = self.__get_choices_internal_get(choice_values, None, wc.key)
for cv in ch_values:
expanded_choice_values.append(
{
**cv,
"weight": float(cv.get("weight", 1.0) * c_weight), # we adjust the weight
}
)
else:
self.warn_or_stop(f"Unsupported choice command '{escape_single_quotes(cmd)}' at {msg_where}!")
else:
expanded_choice_values.append(c)
return expanded_choice_values
def __get_choices_internal_select(
self,
options: dict | None,
choice_values: list[dict],
filter_specifier: Optional[list[list[str]]] = None,
wildcard_key: str = None,
) -> tuple[str, list[str], str, str]:
"""
Select choices based on the options.
Args:
options (dict): The object representing the options construct.
choice_values (list[dict]): A list of choice objects.
filter_specifier (list[list[str]]): The filter specifier.
wildcard_key (str): The wildcard key if it is a wildcard.
Returns:
tuple: A tuple containing the prefix, selected choices, separator and suffix
"""
seen_wildcards_len = len(self.__seen_wildcards)
if options is None:
options = {}
sampler: str = options.get("sampler", "~")
repeating: bool = options.get("repeating", False)
optional: bool = options.get("optional", False)
if "count" in options:
from_value = options["count"]
to_value = from_value
else:
from_value: int = options.get("from", 1)
to_value: int = options.get("to", 1)
separator: str = options.get("separator", self.__state.options.wil_choice_separator)
msg_where = f"wildcard '{escape_single_quotes(wildcard_key)}'" if wildcard_key else "choices"
if sampler != "~":
self.warn_or_stop(f"Unsupported sampler '{escape_single_quotes(sampler)}' at {msg_where} options!")
sampler = "~"
expanded_choice_values = self.__get_choices_internal_get(choice_values, filter_specifier, wildcard_key)
available_choices: list[dict] = []
weights = []
included_choices = 0
excluded_choices = 0
excluded_weights_sum = 0
for i, c in enumerate(expanded_choice_values):
c["choice_index"] = i # we index them to later sort the results
weight = float(c.get("weight", 1.0))
condition = c.get("if", None)
if weight > 0 and (condition is None or self.__eval_condition(condition)):
available_choices.append(c)
weights.append(weight)
included_choices += 1
else:
weights.append(-1)
excluded_choices += 1
excluded_weights_sum += weight
if excluded_choices > 0: # we need to redistribute the excluded weights
weights = [weight + excluded_weights_sum / included_choices for weight in weights if weight >= 0]
weights = np.array(weights)
weights /= weights.sum() # normalize weights
if available_choices:
if from_value < 0:
from_value = 1
elif from_value > len(available_choices):
from_value = len(available_choices)
if to_value < 1:
to_value = 1
elif (to_value > len(available_choices) and not repeating) or from_value > to_value:
to_value = len(available_choices)
num_choices = (
self.__rng.integers(from_value, to_value, endpoint=True) if from_value < to_value else from_value
)
else:
num_choices = 0
if not optional and from_value > 0:
self.warn_or_stop(f"Not enough choices found for {msg_where}!")
if num_choices < 2:
repeating = False
if self.__debug_level == DEBUG_LEVEL.full:
self.__state.logger.debug(
format_output(
f"Selecting {'optional ' if optional else ''}{'repeating ' if repeating else ''}{num_choices} choice"
+ ("s" if num_choices != 1 else "")
+ (f" and separating with '{escape_single_quotes(separator)}'" if num_choices > 1 else "")
)
)
if num_choices > 0:
selected_choices: list[dict] = (
list(self.__rng.choice(available_choices, size=num_choices, p=weights, replace=repeating))
if available_choices
else []
)
if self.__state.options.wil_keep_choices_order:
selected_choices = sorted(selected_choices, key=lambda x: x["choice_index"])
selected_choices_text = []
prefix: str = (
self.__visit(options.get("prefix", None), False, True)
if options.get("prefix", None) is not None
else ""
)
if prefix != "" and re.match(r"\w", prefix[-1]):
prefix += " "
for i, c in enumerate(selected_choices):
t1 = time.monotonic_ns()
choice_content_obj = c.get("content", c.get("text", None))
if isinstance(choice_content_obj, str):
choice_content = choice_content_obj
else:
choice_content = self.__visit(choice_content_obj, False, True)
t2 = time.monotonic_ns()
if self.__debug_level == DEBUG_LEVEL.full:
self.__state.logger.debug(
f"Adding choice {i+1} ({(t2-t1) / 1_000_000_000:.3f} seconds):\n"
+ textwrap.indent(re.sub(r"\n$", "", choice_content), " ")
)
selected_choices_text.append(choice_content)
suffix: str = (
self.__visit(options.get("suffix", None), False, True)
if options.get("suffix", None) is not None
else ""
)
if suffix != "" and re.match(r"\w", suffix[0]):
suffix = " " + suffix
# remove comments
results = [re.sub(r"\s*#[^\n]*(?:\n|$)", "", r, flags=re.DOTALL) for r in selected_choices_text]
else:
prefix = ""
suffix = ""
results = []
if self.__debug_level == DEBUG_LEVEL.full:
list_unseen = [f"'{escape_single_quotes(x)}'" for x in self.__seen_wildcards[seen_wildcards_len:]]
self.__state.logger.debug(f"Unseen wildcards: {', '.join(list_unseen)}")
self.__seen_wildcards = self.__seen_wildcards[:seen_wildcards_len]
return (prefix, results, separator, suffix)
def __get_choices(
self,
options: dict | None,
choice_values: list[dict],
filter_specifier: Optional[list[list[str]]] = None,
wildcard_key: str = None,
) -> str:
r = self.__get_choices_internal_select(options, choice_values, filter_specifier, wildcard_key)
if r[1]:
return r[0] + r[2].join(r[1]) + r[3]
return ""
def __convert_choices_options(self, options: Optional[lark.Tree], is_wcdef: bool = False) -> dict:
"""
Convert the choices options to a dictionary.
Args:
options (Tree): The choices options tree.
Returns:
dict: The converted choices options.
"""
if options is None:
return None
options_dict = {}
if len(options.children) == 1:
if options.children[0] is not None:
options_dict["sampler"] = str(options.children[0])
else:
if options.children[0] is not None:
options_dict["sampler"] = str(options.children[0].children[0])
if options.children[1] is not None:
flags = str(options.children[1].children[0])
options_dict["repeating"] = "r" in flags
options_dict["optional"] = "o" in flags
irange = 2
idesc = 3
isep = 4
if options.children[irange] is not None:
if len(options.children[irange].children) == 1:
if options.children[irange].children[0] is not None:
options_dict["count"] = int(options.children[irange].children[0])
else:
options_dict["from"] = (
int(options.children[irange].children[0])
if options.children[irange].children[0] is not None
else 1
)
options_dict["to"] = (
int(options.children[irange].children[1])
if options.children[irange].children[1] is not None
else 1
)
if is_wcdef:
if options.children[idesc] is not None:
options_dict["description"] = str(options.children[idesc].children[0])[1:-1]
else:
isep -= 1 # only wildcard definition options have a description
if options.children[isep] is not None:
options_dict["separator"] = self.__visit(options.children[isep], False, True)
if not options_dict:
options_dict = None
return options_dict
def __convert_choice(self, choice: lark.Tree) -> dict:
"""
Convert the choice to a dictionary.
Args:
choice (Tree): The choice tree.
Returns:
dict: The converted choice.
"""
choice_dict = {}
choice_dict["command"] = choice.children[0] is not None
c_label_obj = choice.children[1]
choice_dict["labels"] = [str(x).lower() for x in c_label_obj.children[1:-1]] if c_label_obj is not None else []
choice_dict["weight"] = float(choice.children[2].children[0]) if choice.children[2] is not None else 1.0
choice_dict["if"] = choice.children[3].children[0] if choice.children[3] is not None else None
choice_dict["content"] = choice.children[-1]
return choice_dict
def __check_wildcard_initialization(self, wildcard: PPPWildcard) -> tuple[dict | None, list[dict] | None]:
"""
Initializes a wildcard if it hasn't been yet.
Args:
wildcard (PPPWildcard): The wildcard to check.
Returns:
tuple: A tuple containing the options and choice values of the wildcard.
"""
choice_values = wildcard.choices
options = wildcard.options
if choice_values is None:
t1 = time.monotonic_ns()
choice_values = []
options, n = self.get_wildcard_options(wildcard)
# we process the choices
for cv in wildcard.unprocessed_choices[n:]:
if isinstance(cv, dict):
if self.__state.wildcards_obj.is_dict_choice_options(cv):
condition = cv.get("if", None)
if condition is not None and isinstance(condition, str):
try:
cv["if"] = parse_prompt(
self.__state, "condition", condition, self.__state.parsers["condition"], True
)
except lark.exceptions.UnexpectedInput as e:
self.warn_or_stop(
f"Error parsing condition '{escape_single_quotes(condition)}' in wildcard '{escape_single_quotes(wildcard.key)}'! : {e.__class__.__name__}",
e,
)
cv["if"] = None
content = cv.get("content", cv.get("text", None))
cv["content"] = content
if "text" in cv:
del cv["text"]
if content is not None and isinstance(content, str):
try:
cv["content"] = parse_prompt(
self.__state, "choicevalue", content, self.__state.parsers["choicevalue"], True
)
except lark.exceptions.UnexpectedInput as e:
self.warn_or_stop(
f"Error parsing choice content '{escape_single_quotes(content)}' in wildcard '{escape_single_quotes(wildcard.key)}'! : {e.__class__.__name__}",
e,
)
cv["content"] = None
if cv["content"] is not None:
if self.__debug_level == DEBUG_LEVEL.full:
self.__state.logger.debug(f"Processed choice {cv}")
choice_values.append(cv)
else:
self.warn_or_stop(
f"Invalid choice {cv} in wildcard '{escape_single_quotes(wildcard.key)}'!"
)
else:
self.warn_or_stop(f"Invalid choice {cv} in wildcard '{escape_single_quotes(wildcard.key)}'!")
else:
try:
choice_values.append(
self.__convert_choice(
parse_prompt(self.__state, "choice", cv, self.__state.parsers["choice"], True)
)
)
except lark.exceptions.UnexpectedInput as e:
self.warn_or_stop(
f"Error parsing choice '{escape_single_quotes(cv)}' in wildcard '{escape_single_quotes(wildcard.key)}'! : {e.__class__.__name__}",
e,
)
wildcard.choices = choice_values
t2 = time.monotonic_ns()
if self.__debug_level == DEBUG_LEVEL.full:
self.__state.logger.debug(
f"Processed choices for wildcard '{escape_single_quotes(wildcard.key)}' ({(t2-t1) / 1_000_000_000:.3f} seconds)"
)
return (self.__clean_wildcard_options(options), choice_values)
def get_wildcard_options(self, wildcard: PPPWildcard) -> tuple[dict | None, int]:
options = wildcard.options
n = 0
# we check the first choice to see if it is actually options
if isinstance(wildcard.unprocessed_choices[0], dict):
if self.__state.wildcards_obj.is_dict_wcdef_options(wildcard.unprocessed_choices[0]):
options = wildcard.unprocessed_choices[0]
prefix = options.get("prefix", None)
if prefix is not None and isinstance(prefix, str):
try:
options["prefix"] = parse_prompt(
self.__state, "choicevalue", prefix, self.__state.parsers["choicevalue"], True
)
except lark.exceptions.UnexpectedInput as e:
self.warn_or_stop(
f"Error parsing choice prefix '{escape_single_quotes(prefix)}' in wildcard '{escape_single_quotes(wildcard.key)}'! : {e.__class__.__name__}",
e,
)
suffix = options.get("suffix", None)
if suffix is not None and isinstance(suffix, str):
try:
options["suffix"] = parse_prompt(
self.__state, "choicevalue", suffix, self.__state.parsers["choicevalue"], True
)
except lark.exceptions.UnexpectedInput as e:
self.warn_or_stop(
f"Error parsing choice suffix '{escape_single_quotes(suffix)}' in wildcard '{escape_single_quotes(wildcard.key)}'! : {e.__class__.__name__}",
e,
)
n = 1
else:
if wildcard.unprocessed_choices[0].endswith("$$"):
try:
options = self.__convert_choices_options(
parse_prompt(
self.__state,
"as wildcard options",
wildcard.unprocessed_choices[0][:-2].strip(),
self.__state.parsers["wcdefoptions"],
True,
),
True,
)
n = 1
except lark.exceptions.UnexpectedInput:
options = None
if options is None and self.__debug_level == DEBUG_LEVEL.full:
self.__state.logger.debug("Does not have options")
wildcard.options = options
return options, n
def __clean_wildcard_options(self, options: dict) -> dict:
if options is not None:
options.pop("description", None) # description is only for wildcard definitions, not for usage
if not options:
options = None
return options
def wildcard(self, tree: lark.Tree):
"""
Process a wildcard construct in the tree.
"""
t1 = time.monotonic_ns()
seen_wildcards_len = len(self.__seen_wildcards)
start_result = self.result
applied_options = self.__clean_wildcard_options(self.__convert_choices_options(tree.children[0], False))
wildcard_key: str = self.__visit(tree.children[1], False, True)
wc = self.__get_original_node_content(tree, f"?__{wildcard_key}__")
if self.__state.options.wil_process_wildcards:
if self.__debug_level == DEBUG_LEVEL.full:
self.__state.logger.debug(f"Processing wildcard: {wildcard_key}")
selected_wildcards = self.__state.wildcards_obj.get_wildcards(wildcard_key)
if not selected_wildcards:
self.detectedWildcards.append(wc)
self.result += wc
t2 = time.monotonic_ns()
self.__debug_end("wildcard", start_result, t2 - t1, wc)
return
filter_specifier: list[list[str]] = None
filter_object = tree.children[2]
if filter_object is not None:
if ( # it's an inherited filter from another wildcard
isinstance(filter_object.children[1], lark.Token)
and filter_object.children[1] is not None
and "^" in str(filter_object.children[1])
):
filter_wildcard_key = self.__visit(filter_object.children[2], False, True)
filter_specifier = self.__wildcard_filters.get(filter_wildcard_key, None)
if self.__debug_level == DEBUG_LEVEL.full:
self.__state.logger.debug("Filtering choices with inherited filter")
else:
filter_specifier = self.__extract_filter_specifiers(filter_object.children[2])
if self.__debug_level == DEBUG_LEVEL.full:
self.__state.logger.debug("Filtering choices")
self.__wildcard_filters[wildcard_key] = filter_specifier
if filter_object.children[1] is not None and "#" in str(
filter_object.children[1]
): # means do not use the filter in this wildcard
if self.__debug_level == DEBUG_LEVEL.full:
self.__state.logger.debug("Ignoring filter")
filter_specifier = None
else:
filter_specifier = self.__state.wildcards_obj.get_wildcard_default_filter(wildcard_key)
if filter_specifier is not None:
self.__wildcard_filters[wildcard_key] = filter_specifier
if self.__debug_level == DEBUG_LEVEL.full:
self.__state.logger.debug("Applying default filter")
if (
len(selected_wildcards) > 1
and filter_specifier is not None
and any(y[0].isdecimal() for x in filter_specifier for y in x)
):
self.__state.logger.warning(
f"Using a globbing wildcard '{escape_single_quotes(wildcard_key)}' with positional index filters is not recommended!"
)
var_object = tree.children[3]
variablename = None
variablebackup = None
if var_object is not None:
variablename = str(var_object.children[0])
variablevalue = self.__visit(var_object.children[1], False, True)
variablebackup = self.__state.user_variables.get(variablename, None)
self.__remove_user_variable(variablename)
self.__set_user_variable_value(variablename, variablevalue)
choice_values_all = []
for wildcard in selected_wildcards:
if wildcard is None:
self.detectedWildcards.append(wc)
self.result += wc
t2 = time.monotonic_ns()
self.__debug_end("wildcard", start_result, t2 - t1, wc)
return
if wildcard.key in self.__seen_wildcards:
self.warn_or_stop(
f"Circular reference detected with wildcard '{escape_single_quotes(self.__seen_wildcards[-1])}' (chain starts at '{escape_single_quotes(self.__seen_wildcards[0])}')!"
)
continue
self.__seen_wildcards.append(wildcard.key)
if self.__debug_level == DEBUG_LEVEL.full:
self.__state.logger.debug(f"Seen wildcard '{escape_single_quotes(wildcard.key)}'")
(options, choice_values) = self.__check_wildcard_initialization(wildcard)
if options is not None:
if applied_options is None:
applied_options = options
else:
if self.__debug_level == DEBUG_LEVEL.full:
self.__state.logger.debug(
f"Options for wildcard '{escape_single_quotes(wildcard.key)}' are ignored!"
)
choice_values_all += choice_values
self.result += self.__get_choices(applied_options, choice_values_all, filter_specifier, wildcard_key)
if wildcard_key in self.__wildcard_filters:
del self.__wildcard_filters[wildcard_key]
if variablename is not None:
self.__remove_user_variable(variablename)
if variablebackup is not None:
self.__state.user_variables[variablename] = variablebackup
elif self.__state.options.wil_ifwildcards != IFWILDCARDS_CHOICES.remove:
self.detectedWildcards.append(wc)
self.result += wc
if self.__debug_level == DEBUG_LEVEL.full:
list_unseen = [f"'{escape_single_quotes(x)}'" for x in self.__seen_wildcards[seen_wildcards_len:]]
self.__state.logger.debug(f"Unseen wildcards: {', '.join(list_unseen)}")
self.__seen_wildcards = self.__seen_wildcards[:seen_wildcards_len]
t2 = time.monotonic_ns()
self.__debug_end("wildcard", start_result, t2 - t1, f"'{escape_single_quotes(wc)}'")
def __extract_filter_specifiers(self, filters: lark.Tree) -> list[list[str]]:
filter_specifier = []
for or_ in filters.children:
for and_ in or_.children:
label = and_.children[0]
if isinstance(label, lark.Token):
# it's a literal, we can use it directly
filter_specifier.append([str(label)])
else:
# it's a variable, we need to evaluate it
v = self.__visit(label, False, True)
filter_specifier.append([v])
return filter_specifier
def choices(self, tree: lark.Tree):
"""
Process a choices construct in the tree.
"""
t1 = time.monotonic_ns()
start_result = self.result
options = self.__convert_choices_options(tree.children[0], False)
choice_values = [self.__convert_choice(c) for c in tree.children[1::]]
ch = self.__get_original_node_content(tree, "?{...}")
if self.__state.options.wil_process_wildcards:
if self.__debug_level == DEBUG_LEVEL.full:
self.__state.logger.debug("Processing choices:")
self.result += self.__get_choices(options, choice_values)
elif self.__state.options.wil_ifwildcards != IFWILDCARDS_CHOICES.remove:
self.detectedWildcards.append(ch)
self.result += ch
t2 = time.monotonic_ns()
self.__debug_end("choices", start_result, t2 - t1, f"'{escape_single_quotes(ch)}'")
def __default__(self, tree):
t1 = time.monotonic_ns()
start_result = self.result
self.__visit(tree.children)
t2 = time.monotonic_ns()
self.__debug_end(tree.data.value, start_result, t2 - t1)
def start(self, tree):
self.result = ""
t1 = time.monotonic_ns()
self.__visit(tree.children)
attention_processing = self.__state.host_config.get("attention", "ok")
# process the found negative tags
for negtag in self.__negtags:
if self.__state.options.cup_mergeattention:
# join consecutive attention elements
for i in range(len(negtag.shell) - 1, 0, -1):
if negtag.shell[i].type == "at" and negtag.shell[i - 1].type == "at":
new_weight = ( # we limit the new weight to two decimals
math.floor(100 * float(negtag.shell[i - 1].data[1]) * float(negtag.shell[i].data[1])) / 100
)
new_weight_str = f"{new_weight:.2f}".rstrip("0").rstrip(".")
if new_weight_str == "0.9" and attention_processing != "parentheses":
new_kind = 1
elif new_weight_str == "1.1":
new_kind = 2
else:
new_kind = 3
negtag.shell[i - 1] = self.AccumulatedShell(
"at",
(new_kind, new_weight_str),
)
negtag.shell.pop(i)
start = ""
end = ""
for s in negtag.shell:
match s.type:
case "at":
if s.data[0] == 1:
start += "["
end = "]" + end
elif s.data[0] == 2:
start += "("
end = ")" + end
else: # 3
start += "("
end = f":{s.data[1]})" + end
# case "sc":
case "scb":
start += "["
end = f"::{s.data}]" + end
case "sca":
start += "["
end = f":{s.data}]" + end
# case "al":
case "alo":
start += "[" + ("|" * int(s.data["pos"] - 1))
end = ("|" * int(s.data["len"] - s.data["pos"])) + "]" + end
content = start + negtag.content + end
position = negtag.parameters or "s"
if position.startswith("i"):
n = int(position[1])
self.insertion_at[n] = [negtag.start, negtag.end]
elif len(content) > 0:
if content not in self.__already_processed:
if self.__state.options.stn_ignore_repeats:
self.__already_processed.append(content)
if self.__debug_level == DEBUG_LEVEL.full:
self.__state.logger.debug(format_output(f"Adding content at position {position}: {content}"))
if position == "e":
self.add_at["end"].append(content)
elif position.startswith("p"):
n = int(position[1])
self.add_at["insertion_point"][n].append(content)
else: # position == "s" or invalid
self.add_at["start"].append(content)
else:
self.__state.logger.warning(format_output(f"Ignoring repeated content: {content}"))
t2 = time.monotonic_ns()
self.__debug_end("start", "", t2 - t1)