from collections import namedtuple import logging 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, log from ppp_utils import escape_single_quotes 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 log(self, kind, message: str, min_level: DEBUG_LEVEL | None = None): log(self.state.logger, self.state.options.debug_level, kind, message, min_level) 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 self.log(logging.INFO, f"Processing {prompt_description}...") self.visit(parsed_prompt) t2 = time.monotonic_ns() self.log(logging.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 # self.log(logging.DEBUG, f"Visiting node {node}.") if restore_state: # self.log(logging.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: # self.log(logging.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 __idxsep_to_idx_sep(self, idxsep: str | None) -> tuple[Optional[int], Optional[str]]: """ Convert an index or separator string to an index and a separator. Args: idxsep (str|None): The index or separator string. Returns: tuple: A tuple containing the index and separator. """ if idxsep is None: return None, None is_quoted = idxsep.startswith(("'", '"')) and idxsep.endswith(("'", '"')) if not idxsep.isdigit() and not is_quoted: # bare identifier: resolve as variable idxsep = self.get_final_user_variable(idxsep) is_quoted = False if idxsep.isdigit(): return int(idxsep), None # separator: strip surrounding quotes if present return None, idxsep[1:-1] if is_quoted else idxsep def __get_user_variable_value( self, name: str, idxsep: str | None = None, evaluate=True, visit=False ) -> str | list[str] | None: """ Get the value of a user variable. Args: name (str): The name of the user variable. idxsep (str|None): The index for an array variable or a separator. evaluate (bool): Whether to evaluate the variable. visit (bool): Whether to also visit the variable (add to result). Returns: str|list[str]|None: The value of the user variable. """ def visit_value(v): visited = False if isinstance(v, lark.Tree): if evaluate: v = self.__visit(v, restore_state=not visit, discard_content=not visit) visited = visit else: v = self.__get_original_node_content(v, "") if visit and not visited: self.result += v return v v = self.state.user_variables.get(name, None) if v is None: return None is_array = name[-2:] == "[]" if is_array: if isinstance(v, list): idx, sep = self.__idxsep_to_idx_sep(idxsep) if idx is not None: if 0 <= idx < len(v): v = visit_value(v[idx]) else: v = None # invalid index else: if sep is None: sep = self.state.options.choice_separator v2 = [] for i, item in enumerate(v): v2.append(visit_value(item)) if visit and i < len(v) - 1: self.result += sep v = v2 else: v = None # error elif idxsep is not None: v = None # error else: v = visit_value(v) return v def __separate_arrayref(self, nameidx: str): """ Separate the name and index/sep part of a variable reference. Args: nameidx (str): The variable reference string. Returns: tuple: A tuple containing the name and index/sep part. """ name, isarray, idxsep = re.match(r"^([^\[]+)(\[([^[]*)\])?$", nameidx).groups() if isarray: name += "[]" if idxsep == "": idxsep = None return name, idxsep def get_final_user_variable(self, nameidx: str) -> str: """ Get the final value of a user variable, resolving any references if needed. Args: nameidx (str): The variable reference string. Returns: str: The final value of the user variable. """ name, idxsep = self.__separate_arrayref(nameidx) v = self.__get_user_variable_value(name, idxsep, True, False) if isinstance(v, list): _, sep = self.__idxsep_to_idx_sep(idxsep) if sep is None: sep = self.state.options.choice_separator v = sep.join(str(item) for item in v) return v def __set_user_variable_value(self, name: str, value: str | lark.Tree | list): """ Set the value of a user variable. Args: name (str): The name of the user variable. value (str|lark.Tree|list): 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.log(logging.DEBUG, f"TreeProcessor.{construct} {info}({duration / 1_000_000_000:.3f} seconds){output}") def __adjust_strnum(self, s: str) -> str | int: """ Adjust a string that may represent a number to its appropriate type. If it is a digit and does not have leading zeros (unless it's just "0"), it is converted to an integer. Args: s (str): The string to adjust. Returns: str | int: The adjusted string or integer. """ if s.isdigit() and not (s.startswith("0") and len(s) > 1): return int(s) return s.lower() def warn_mixedtype(self, desc: str, operand1, operand2, operation): """ Warn the user if mixed type values are used in a comparison. Args: desc (str): Description of the comparison. operand1: The first operand. operand2: The second operand. operation: The operation to perform if the types are compatible. """ if operand1 is None or operand2 is None: self.warn_or_stop(f"Undefined value used in comparison: '{escape_single_quotes(desc)}'") return False if ( isinstance(operand1, (str, int, bool)) and isinstance(operand2, (str, int, bool)) and operand1.__class__ != operand2.__class__ ): self.warn_or_stop(f"Mixed type values used in comparison: '{escape_single_quotes(desc)}'") return False return operation(operand1, operand2) def __resolve_operand(self, c: str) -> str | bool | int: """ Resolve an operand value. Args: c (str): The operand value to resolve. Returns: str | bool | int: The resolved operand value (in lowercase for strings). """ if c.startswith('"') and c.endswith('"') or c.startswith("'") and c.endswith("'"): return self.__adjust_strnum(c[1:-1]) if c.isdigit(): return int(c) if c.lower() in ("false", ""): return False if c.lower() == "true": return True # Bare identifier - resolve as variable reference if c.startswith("_"): vartype = "system" val = self.state.system_variables.get(c, None) else: vartype = "user" varname, varidxsep = self.__separate_arrayref(c) val = self.__get_user_variable_value(varname, varidxsep) if val is None: val = "" self.warn_or_stop(f"Unknown {vartype} variable '{escape_single_quotes(c)}'") if isinstance(val, str): if val.isdigit(): val = int(val) else: val = self.__adjust_strnum(val) if val in ("false", ""): return False if val == "true": return True return val def __eval_basiccondition( self, operand1: str | list[str], operator: str, operand2: str | list[str], ) -> bool: """ Evaluate a condition based on the given operands and operator. Args: operand1 (str | list[str]): The first operand. operator (str): The operator. operand2 (str | list[str]): The second operand. Returns: bool: The result of the condition evaluation. """ if isinstance(operand1, list): operand1_value = list(self.__resolve_operand(c) for c in operand1) else: operand1_value = self.__resolve_operand(operand1) operand1_isarray = isinstance(operand1_value, list) if isinstance(operand2, list): operand2_value = list(self.__resolve_operand(c) for c in operand2) else: operand2_value = self.__resolve_operand(operand2) operand2_isarray = isinstance(operand2_value, list) if operator == "truthy": result = bool(operand1_value) else: op1_desc = f"[{', '.join(operand1)}]" if operand1_isarray else operand1 op2_desc = f"[{', '.join(operand2)}]" if operand2_isarray else operand2 condition_desc = f"{op1_desc} {operator} {op2_desc}" def pairwise_all(op): return len(operand1_value) == len(operand2_value) and all( op(a, b) for a, b in zip(operand1_value, operand2_value) ) def wmt(a, b, op): return self.warn_mixedtype(condition_desc, a, b, op) if not operand1_isarray and not operand2_isarray: operations = { "eq": lambda: wmt(operand1_value, operand2_value, lambda x, y: x == y), "ne": lambda: wmt(operand1_value, operand2_value, lambda x, y: x != y), "gt": lambda: wmt(operand1_value, operand2_value, lambda x, y: x > y), "lt": lambda: wmt(operand1_value, operand2_value, lambda x, y: x < y), "ge": lambda: wmt(operand1_value, operand2_value, lambda x, y: x >= y), "le": lambda: wmt(operand1_value, operand2_value, lambda x, y: x <= y), "in": lambda: wmt(str(operand1_value), str(operand2_value), lambda x, y: x in y), "contains": lambda: wmt(str(operand1_value), str(operand2_value), lambda x, y: y in x), } elif operand1_isarray and operand2_isarray: operations = { "eq": lambda: pairwise_all(lambda a, b: wmt(a, b, lambda x, y: x == y)), "ne": lambda: pairwise_all(lambda a, b: wmt(a, b, lambda x, y: x != y)), "gt": lambda: pairwise_all(lambda a, b: wmt(a, b, lambda x, y: x > y)), "lt": lambda: pairwise_all(lambda a, b: wmt(a, b, lambda x, y: x < y)), "ge": lambda: pairwise_all(lambda a, b: wmt(a, b, lambda x, y: x >= y)), "le": lambda: pairwise_all(lambda a, b: wmt(a, b, lambda x, y: x <= y)), "in": lambda: all(wmt(a, operand2_value, lambda x, y: x in y) for a in operand1_value), "contains": lambda: all(wmt(a, operand1_value, lambda x, y: x in y) for a in operand2_value), } elif operand1_isarray and not operand2_isarray: operations = { "eq": lambda: False, "ne": lambda: True, # "gt": lambda: False, # "lt": lambda: False, # "ge": lambda: False, # "le": lambda: False, "contains": lambda: wmt(operand1_value, operand2_value, lambda x, y: y in x), } elif not operand1_isarray and operand2_isarray: operations = { "eq": lambda: False, "ne": lambda: True, # "gt": lambda: False, # "lt": lambda: False, # "ge": lambda: False, # "le": lambda: False, "in": lambda: wmt(operand1_value, operand2_value, lambda x, y: x in y), } else: operations = {} if operator not in operations: self.warn_or_stop( f"Unsupported operator '{escape_single_quotes(operator)}' in condition '{escape_single_quotes(condition_desc)}'" ) return False result = operations[operator]() return result def __separate_vardescriptor(self, vardescriptor: lark.Tree) -> tuple[str, str | None]: """ Separate the name and index/sep part of a variable descriptor. Args: vardescriptor (lark.Tree): The variable descriptor tree. Returns: tuple[str, str | None]: A tuple containing the name and index/sep part of the variable descriptor. """ vardescriptor_name = str(vardescriptor.children[0]) vardescriptor_idx = None if vardescriptor.children[1] is not None: vardescriptor_name += "[]" if vardescriptor.children[2] is not None: vardescriptor_idx = vardescriptor.children[2] return vardescriptor_name, vardescriptor_idx def __get_complex_element(self, value_node: lark.Tree | lark.Token) -> str: """ Return in string form the value of a complex value node, which can be either a simple value or a variable descriptor. Does not evaluate variables. Args: value_node (lark.Tree | lark.Token): The complex value node to be evaluated. Returns: str: The value. """ if isinstance(value_node, lark.Tree): # it's a vardescriptor_get vardescriptor_name, vardescriptor_idx = self.__separate_vardescriptor(value_node) return ( vardescriptor_name[0:-2] + f"[{vardescriptor_idx}]" if vardescriptor_idx is not None else vardescriptor_name ) # it's a SIMPLEVALUE return str(value_node) def __get_cond_operand(self, value_node: lark.Tree) -> str | list[str]: """ Returns a simple value or a list of simple values or variables. Does not evaluate them. Args: value_node (lark.Tree): The value tree to be evaluated. Returns: str or list[str]: The result. """ if isinstance(value_node, lark.Tree) and value_node.data == "listvalue": return list(self.__get_complex_element(v) for v in value_node.children) return self.__get_complex_element(value_node) def __eval_condition(self, condition: lark.Tree) -> bool: """ Evaluate an if condition based on the given condition tree. Args: condition (lark.Tree): The condition tree to be evaluated. Returns: bool: The result of the if condition evaluation. """ # self.log(logging.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 # we get the name of the variable cond_operand1 = self.__get_cond_operand(condition.children[0]) poscomp = 1 invert = False if poscomp >= len(condition.children): # no condition, just a variable cond_operation = "truthy" cond_operand2 = "true" else: # we get the comparison (with possible not) and the value cond_operation = str(condition.children[poscomp]) if cond_operation == "not": invert = not invert poscomp += 1 cond_operation = str(condition.children[poscomp]) poscomp += 1 cond_value_node = condition.children[poscomp] cond_operand2 = self.__get_cond_operand(cond_value_node) cond_result = self.__eval_basiccondition(cond_operand1, cond_operation, cond_operand2) 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() ) self.log(logging.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": self.log(logging.DEBUG, "Scheduling construct removed, taking before option") if before is not None: self.__visit(before) elif scheduling_processing == "after": self.log(logging.DEBUG, "Scheduling construct removed, taking after option") if after is not None: self.__visit(after) elif scheduling_processing == "first": self.log(logging.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": self.log(logging.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: self.log(logging.DEBUG, f"Shell scheduled before with position {pos}") self.__shell.append(self.AccumulatedShell("scb", pos)) self.__visit(before) self.__shell.pop() self.log(logging.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_empty_constructs 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": self.log(logging.DEBUG, "Alternation construct removed, taking first option") self.__visit(tree.children[0]) elif alternation_processing == "remove": self.log(logging.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): self.log(logging.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_empty_constructs 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" self.log(logging.DEBUG, f"Shell attention with weight {weight}") current_tree = tree.children[0] if self.state.options.cup_merge_attention: 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" self.log(logging.DEBUG, "Converted to parentheses format") elif attention_processing == "disable": weight_kind = 0 self.log(logging.DEBUG, "Attention construct disabled") elif attention_processing == "remove": weight_kind = -1 self.log(logging.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_empty_constructs 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_name: str, variable_idx: str | None, 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_name.startswith("_"): self.warn_or_stop( f"Invalid variable name '{escape_single_quotes(variable_name)}' detected! System variables cannot be set." ) return info = variable_name is_array = variable_name[-2:] == "[]" if variable_idx is not None: if not is_array: self.warn_or_stop( f"Invalid variable name '{escape_single_quotes(variable_name)}' detected! Only array variables can be indexed." ) return info = f"{variable_name[0:-2]}[{variable_idx}]" value_description = self.__get_original_node_content(content, None) value = content raw_oldvalue = self.state.user_variables.get(variable_name, None) newvalue = None some_error = False if variable_idx is not None: if ( is_array and raw_oldvalue is not None and isinstance(raw_oldvalue, list) and not 0 <= int(variable_idx) < len(raw_oldvalue) ): self.warn_or_stop( f"Invalid index {variable_idx} for variable '{escape_single_quotes(variable_name)}'! Index out of bounds." ) some_error = True if not is_array: self.warn_or_stop( f"Invalid index for '{escape_single_quotes(variable_name)}'! Only array variables can be indexed." ) some_error = True if raw_oldvalue is None or not isinstance(raw_oldvalue, list): self.warn_or_stop( f"Invalid index for '{escape_single_quotes(variable_name)}'! Cannot set index of an undefined or non-array variable." ) some_error = True adding = False if not some_error: 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"]): adding = True info += f" += '{escape_single_quotes(value_description or '')}'" if raw_oldvalue is None: newvalue = value self.warn_or_stop(f"Unknown variable {escape_single_quotes(variable_name)}") elif is_array: if not isinstance(raw_oldvalue, list): self.warn_or_stop( f"Invalid variable value for '{escape_single_quotes(variable_name)}'! Cannot add to a non-array value." ) else: newvalue = value 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: # lark.Tree 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 '')}'" if raw_oldvalue is None: newvalue = value else: info += " (not set)" else: newvalue = value if newvalue is not None: is_starred = isinstance(newvalue, lark.Tree) and newvalue.data == "starredvalue" access_full_array = is_array and variable_idx is None if any(item in modifiers_str for item in ["!", "evaluate"]): newvalue = self.__visit(newvalue, False, True) info += " =! " else: info += " = " if is_starred: if access_full_array and isinstance(newvalue.children[0], lark.Tree): if newvalue.children[0].data == "vardescriptor_get": vardescriptor_name, vardescriptor_idx = self.__separate_vardescriptor(newvalue.children[0]) if vardescriptor_idx is not None: newvalue = None else: newvalue = self.__get_user_variable_value(vardescriptor_name, None) elif newvalue.children[0].data == "listvalue": newvalue = list( self.__resolve_operand(c) for c in self.__get_cond_operand(newvalue.children[0]) ) elif newvalue.children[0].data == "wildcard": backup_result = self.result newvalue = self.__process_wildcard(newvalue.children[0]) self.result = backup_result else: newvalue = None else: newvalue = None if newvalue is None: self.warn_or_stop( f"Invalid use of starred value for '{escape_single_quotes(variable_name)}'! Starred values can only be assigned or added to unindexed array variables." ) newvalue = "" if is_array: if variable_idx is not None: # Accessing an existing index, we need to update the array newarray = raw_oldvalue.copy() newarray[int(variable_idx)] = newvalue newvalue = newarray else: # Accessing the whole array if isinstance(newvalue, list): if raw_oldvalue is None or not adding: pass else: newvalue = raw_oldvalue + newvalue else: if raw_oldvalue is None or not adding: newvalue = [newvalue] else: newvalue = raw_oldvalue + [newvalue] self.__set_user_variable_value(variable_name, newvalue) currentvalue = self.__get_user_variable_value(variable_name, variable_idx, False) if currentvalue is None: info += "error" elif isinstance(currentvalue, list): info += "[" + ", ".join(f"'{escape_single_quotes(str(v))}'" for v in currentvalue) + "]" 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) vardescriptor_name, vardescriptor_idx = self.__separate_vardescriptor(tree.children[0]) self.__varset("variableset", vardescriptor_name, vardescriptor_idx, 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. """ vardescriptor_name, vardescriptor_idx = self.__separate_vardescriptor(tree.children[0]) self.__varset("commandset", vardescriptor_name, vardescriptor_idx, tree.children[1], tree.children[2]) def __varecho( self, command: str, variable_name: str, variable_idxsep: str | None, 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 is_array = variable_name[-2:] == "[]" vname = f"{variable_name[0:-2]}[{variable_idxsep}]" if variable_idxsep is not None else variable_name value = self.__get_user_variable_value(variable_name, variable_idxsep, True, True) if value is None: if default is not None: self.log(logging.DEBUG, f"Variable '{escape_single_quotes(vname)}' not found, using default value") v = self.__visit(default, False, True) self.result += v default_value = v self.state.echoed_variables[vname] = v else: self.warn_or_stop(f"Unknown variable {escape_single_quotes(vname)}") default_value = "" self.state.echoed_variables[vname] = "" else: self.state.echoed_variables[vname] = value t2 = time.monotonic_ns() info = variable_name if is_array and variable_idxsep is not None: info = info[0:-2] + f"[{variable_idxsep}]" 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. """ vardescriptor_name, vardescriptor_idxsep = self.__separate_vardescriptor(tree.children[0]) self.__varecho( "variableuse", vardescriptor_name, vardescriptor_idxsep, tree.children[1] if len(tree.children) > 1 else None, ) def commandecho(self, tree: lark.Tree): """ Process an echo command in the tree. """ vardescriptor_name, vardescriptor_idxsep = self.__separate_vardescriptor(tree.children[0]) self.__varecho( "commandecho", vardescriptor_name, vardescriptor_idxsep, 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.cup_remove_extranetwork_tags: 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: self.log( logging.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: self.log( logging.INFO, f"Mapping extranetwork '{escape_single_quotes(extnet_id)}' to just triggers", ) extnet_id = None else: if not found_in_cache: self.log( logging.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_extranetwork_tags: 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: self.log(logging.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: self.log(logging.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.cup_remove_extranetwork_tags: 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"Included wildcard '{escape_single_quotes(cmd_args)}' not found 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) self.log(logging.DEBUG, f"Seen wildcard '{escape_single_quotes(wc.key)}'") self.log(logging.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_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.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 self.log( logging.DEBUG, 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.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() self.log( logging.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 = [] self.log( logging.DEBUG, "Unseen wildcards: " + ", ".join([f"'{escape_single_quotes(x)}'" for x in self.__seen_wildcards[seen_wildcards_len:]]), ) self.__seen_wildcards = self.__seen_wildcards[:seen_wildcards_len] return (prefix, results, separator, suffix) 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: self.log(logging.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() self.log( logging.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: self.log(logging.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 __process_wildcard(self, tree: lark.Tree) -> list: """ Process a wildcard in the tree. Returns: list: A list containing the selected elements. """ t1 = time.monotonic_ns() chosen_choices = [] start_result = self.result seen_wildcards_len = len(self.__seen_wildcards) 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.process_wildcards: self.log(logging.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) self.log(logging.DEBUG, "Filtering choices with inherited filter") else: filter_specifier = self.__extract_filter_specifiers(filter_object.children[2]) self.log(logging.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 self.log(logging.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 self.log(logging.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.log( logging.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: vardescriptor_name, vardescriptor_idx = self.__separate_vardescriptor(var_object.children[0]) variablename = vardescriptor_name # 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) self.__varset("wildcard", variablename, vardescriptor_idx, None, var_object.children[1]) 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) self.log(logging.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: self.log( logging.DEBUG, f"Options for wildcard '{escape_single_quotes(wildcard.key)}' are ignored!" ) choice_values_all += choice_values prefix, chosen_choices, separator, suffix = self.__get_choices_select( applied_options, choice_values_all, filter_specifier, wildcard_key ) if chosen_choices: self.result += prefix + separator.join(chosen_choices) + suffix 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.if_wildcards != 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.log(logging.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)}'") return chosen_choices def wildcard(self, tree: lark.Tree): """ Process a wildcard construct in the tree. """ self.__process_wildcard(tree) 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.process_wildcards: self.log(logging.DEBUG, "Processing choices:") prefix, chosen_choices, separator, suffix = self.__get_choices_select(options, choice_values) if chosen_choices: self.result += prefix + separator.join(chosen_choices) + suffix elif self.state.options.if_wildcards != 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_merge_attention: # 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) self.log(logging.DEBUG, 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.log(logging.WARNING, f"Ignoring repeated content: {content}") t2 = time.monotonic_ns() self.__debug_end("start", "", t2 - t1)