2483 lines
118 KiB
Python
2483 lines
118 KiB
Python
from collections import namedtuple
|
|
from functools import reduce
|
|
from itertools import combinations, combinations_with_replacement, permutations, product
|
|
import logging
|
|
import math
|
|
import re
|
|
import textwrap
|
|
import time
|
|
from typing import Any, 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.
|
|
"""
|
|
|
|
NEGATIVE_SEP = "\x1d"
|
|
AccumulatedShell = namedtuple("AccumulatedShell", ["type", "data"])
|
|
NegTag = namedtuple("NegTag", ["start", "end", "content", "parameters", "shell"])
|
|
|
|
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.__shell: list[TreeProcessor.AccumulatedShell] = [] # type: ignore
|
|
self.__negtags: list[TreeProcessor.NegTag] = [] # type: ignore
|
|
self.__already_processed: list[str] = []
|
|
self.__is_negative = False
|
|
self.__wildcard_filters = {}
|
|
self.__seen_wildcards: list[str] = []
|
|
self.__add_at: dict[str, list] = {"start": [], "insertion_point": [[] for _ in range(10)], "end": []}
|
|
self.__insertion_at: list[tuple[int, int]] = [None for _ in range(10)]
|
|
self.__detectedWildcards: list[tuple[str, bool]] = []
|
|
self.__result = ""
|
|
self.__comb_forced_path: list[int] = []
|
|
self.__comb_trace: list[int] = []
|
|
self.__cycl_forced_path: list[int] = []
|
|
self.__cycl_trace: list[int] = []
|
|
|
|
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 __reset_run_state(self):
|
|
"""Reset all per-run mutable state for a fresh combinatorial pass."""
|
|
self.__shell = []
|
|
self.__negtags = []
|
|
self.__already_processed = []
|
|
self.__is_negative = False
|
|
self.__wildcard_filters = {}
|
|
self.__seen_wildcards = []
|
|
self.__add_at = {"start": [], "insertion_point": [[] for _ in range(10)], "end": []}
|
|
self.__insertion_at = [None for _ in range(10)]
|
|
self.__detectedWildcards = []
|
|
self.__result = ""
|
|
if self.state.extranetwork_mappings_obj is not None:
|
|
self.state.extranetwork_mappings_obj.cached_mappings.clear()
|
|
|
|
def start_visit(
|
|
self,
|
|
parsed: lark.Tree,
|
|
) -> list[tuple[str, list[tuple[str, bool]], tuple[dict[str, Any], dict[str, str]]]]:
|
|
"""
|
|
Process the positive and negative prompts in a unified way using the same processor.
|
|
STN insertions are applied to the negative result directly inside this processor.
|
|
|
|
Args:
|
|
parsed (Tree): The parsed unified prompt.
|
|
|
|
Returns:
|
|
list[tuple[str, list[tuple[str,bool]], tuple[dict[str, Any], dict[str, str]]]]: A list of
|
|
(processed prompt, detected wildcards, variables snapshot) triples — one entry per
|
|
combination in combinatorial mode, or a single entry otherwise. The variables snapshot
|
|
is the return value of ``state.variables.backup_user_and_echoed()`` captured after processing.
|
|
"""
|
|
self.log(logging.INFO, "Processing prompt...")
|
|
|
|
self.__detectedWildcards = []
|
|
self.__is_negative = False
|
|
self.__result = ""
|
|
|
|
if not self.state.options.do_combinatorial:
|
|
self.__cycl_forced_path = list(self.state.cyclical_state.current_path)
|
|
self.__cycl_trace = []
|
|
self.visit(parsed)
|
|
self.__finalize_echoed_variables()
|
|
if self.__cycl_trace:
|
|
self.state.cyclical_state.last_trace = self.__cycl_trace[:]
|
|
self.state.cyclical_state.advance()
|
|
return [(self.__result, self.__detectedWildcards, self.state.variables.backup_user_and_echoed())]
|
|
|
|
# Combinatorial mode: explore every possible path through choices and wildcards via DFS.
|
|
# __comb_forced_path drives which option is selected at each decision point;
|
|
# __comb_trace records how many options were available at each point so the DFS can
|
|
# correctly enumerate unexplored branches after each run.
|
|
initial_vars = self.state.variables.backup_user_and_echoed()
|
|
results: list[tuple[str, list[tuple[str, bool]], tuple]] = []
|
|
limit = self.state.options.combinatorial_limit
|
|
|
|
def _run(forced_path: tuple[int, ...]) -> tuple[int, ...]:
|
|
self.log(logging.DEBUG, f"Running combinatorial path: {forced_path}")
|
|
self.__comb_forced_path = list(forced_path)
|
|
self.__comb_trace = []
|
|
self.__reset_run_state()
|
|
self.state.variables.restore_user_and_echoed(initial_vars)
|
|
self.visit(parsed)
|
|
self.__finalize_echoed_variables()
|
|
results.append(
|
|
(self.__result, self.__detectedWildcards.copy(), self.state.variables.backup_user_and_echoed())
|
|
)
|
|
if len(results) == 1:
|
|
first_run_estimate = reduce(lambda x, y: x * y, self.__comb_trace, 1)
|
|
self.log(logging.INFO, f"Estimated combinations (lower bound): {first_run_estimate}")
|
|
self.log(logging.INFO, f"Added combination {len(results)}")
|
|
return tuple(self.__comb_trace)
|
|
|
|
limit_reached = False
|
|
|
|
def _dfs(forced_path: tuple[int, ...]):
|
|
nonlocal limit_reached
|
|
if 0 < limit <= len(results):
|
|
limit_reached = True
|
|
return
|
|
trace = _run(forced_path)
|
|
# For each decision that was reached but not forced, spawn branches for all
|
|
# options beyond the default (index 0).
|
|
# Iterate in reverse so later decisions vary fastest, producing lexicographic order.
|
|
for i in range(len(trace) - 1, len(forced_path) - 1, -1):
|
|
if 0 < limit <= len(results):
|
|
limit_reached = True
|
|
return
|
|
num_options = trace[i]
|
|
for opt in range(1, num_options):
|
|
if 0 < limit <= len(results):
|
|
limit_reached = True
|
|
return
|
|
# Pad with zeros for intermediate decisions so they keep the default.
|
|
new_path = forced_path + (0,) * (i - len(forced_path)) + (opt,)
|
|
_dfs(new_path)
|
|
|
|
_dfs(())
|
|
if limit_reached:
|
|
self.log(logging.WARNING, f"Combinatorial limit of {limit} reached; some combinations have been skipped.")
|
|
return results
|
|
|
|
def __finalize_echoed_variables(self):
|
|
var_keys = self.state.variables.all_user_or_echoed_keys()
|
|
for k in var_keys:
|
|
ev = self.state.variables.get_echoed_value(k)
|
|
if ev is None:
|
|
ev = self.state.variables.get_user(k)
|
|
if ev is None or ev.__class__ != str: # strict check to avoid problems with Tokens
|
|
self.log(logging.DEBUG, f"Completing variable: {k}")
|
|
ev = self.get_final_variable(k)
|
|
self.state.variables.echo(k, ev) # ensure all variables are echoed so they are included in the snapshot
|
|
|
|
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_vars = self.state.variables.backup_user_and_echoed()
|
|
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.variables.restore_user_and_echoed(backup_vars)
|
|
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 __parse_array_specifier(self, specifier: str | None) -> tuple[Optional[int], Optional[str], Optional[bool]]:
|
|
"""
|
|
Convert an index/separator/count string to a specific type.
|
|
|
|
Args:
|
|
specifier (str|None): The specifier string.
|
|
|
|
Returns:
|
|
tuple: A tuple containing the index, separator and count boolean.
|
|
"""
|
|
if specifier is None:
|
|
return None, None, None
|
|
if specifier == "#": # special value to indicate length of the array variable
|
|
return None, None, True
|
|
if specifier.startswith("&"): # special value to indicate a separator
|
|
return None, specifier[2:-1], False
|
|
if not specifier.isdecimal():
|
|
# bare identifier: resolve as variable
|
|
specifier = self.get_final_variable(specifier)
|
|
if specifier.isdecimal():
|
|
return int(specifier), None, False
|
|
# invalid specifier
|
|
return None, None, False
|
|
|
|
def __get_variable_value(
|
|
self, name: str, specifier: str | None = None, evaluate=True, visit=False
|
|
) -> str | list[str] | None:
|
|
"""
|
|
Get the value of a variable.
|
|
|
|
Args:
|
|
name (str): The name of the variable.
|
|
specifier (str|None): The specifier for an array variable.
|
|
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 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, "")
|
|
elif isinstance(v, lark.Token):
|
|
v = str(v)
|
|
if visit and not visited:
|
|
self.__result += v
|
|
return v
|
|
|
|
v = self.state.variables.get(name)
|
|
if v is None:
|
|
return None
|
|
is_array = name[-2:] == "[]"
|
|
if is_array:
|
|
if isinstance(v, list):
|
|
idx, sep, cnt = self.__parse_array_specifier(specifier)
|
|
if cnt:
|
|
v = len(v)
|
|
if visit:
|
|
self.__result += str(v)
|
|
elif 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 specifier is not None:
|
|
v = None # error
|
|
else:
|
|
v = visit_value(v)
|
|
return v
|
|
|
|
def __separate_arrayref(self, name_specifier: str):
|
|
"""
|
|
Separate the name and specifier part of a variable reference.
|
|
|
|
Args:
|
|
name_specifier (str): The variable reference string.
|
|
|
|
Returns:
|
|
tuple: A tuple containing the name and specifier part.
|
|
"""
|
|
name, isarray, specifier = re.match(r"^([^\[]+)(\[([^[]*)\])?$", name_specifier).groups()
|
|
if isarray:
|
|
name += "[]"
|
|
if specifier == "":
|
|
specifier = None
|
|
return name, specifier
|
|
|
|
def get_final_variable(self, name_specifier: str) -> str:
|
|
"""
|
|
Get the final value of a variable, resolving any references if needed.
|
|
|
|
Args:
|
|
name_specifier (str): The variable reference string.
|
|
|
|
Returns:
|
|
str: The final value of the variable.
|
|
"""
|
|
name, specifier = self.__separate_arrayref(name_specifier)
|
|
v = self.__get_variable_value(name, specifier, True, False)
|
|
if isinstance(v, list):
|
|
_, sep, _ = self.__parse_array_specifier(specifier)
|
|
if sep is None:
|
|
sep = self.state.options.choice_separator
|
|
v = sep.join(str(item) for item in v)
|
|
return str(v)
|
|
|
|
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 | float:
|
|
"""
|
|
Adjust a string that may represent a number to its appropriate type.
|
|
If it is a number, it is converted to an integer or float.
|
|
|
|
Args:
|
|
s (str): The string to adjust.
|
|
|
|
Returns:
|
|
str | int | float: The adjusted string, integer, or float.
|
|
"""
|
|
try:
|
|
return int(s)
|
|
except ValueError:
|
|
pass
|
|
if bool(re.match(r"^[+-]?\d+\.\d+$", s)):
|
|
return float(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
|
|
compatible_types = [
|
|
(str, str),
|
|
(int, int),
|
|
(float, float),
|
|
(bool, bool),
|
|
(int, float),
|
|
(float, int),
|
|
(str, list),
|
|
(int, list),
|
|
(float, list),
|
|
(bool, list),
|
|
]
|
|
if not any(isinstance(operand1, t1) and isinstance(operand2, t2) for t1, t2 in compatible_types):
|
|
self.warn_or_stop(
|
|
f"Mixed type values ({type(operand1).__name__}, {type(operand2).__name__}) used in comparison: '{escape_single_quotes(desc)}'"
|
|
)
|
|
return False
|
|
return operation(operand1, operand2)
|
|
|
|
def __resolve_operand(self, c: str) -> str | bool | int | float:
|
|
"""
|
|
Resolve an operand value.
|
|
|
|
Args:
|
|
c (str): The operand value to resolve.
|
|
|
|
Returns:
|
|
str | bool | int | float: The resolved operand value (in lowercase for strings).
|
|
"""
|
|
if c.startswith('"') and c.endswith('"') or c.startswith("'") and c.endswith("'"):
|
|
return c[1:-1].lower() if self.state.options.strict_operators else self.__adjust_strnum(c[1:-1])
|
|
try:
|
|
return int(c)
|
|
except ValueError:
|
|
pass
|
|
if bool(re.match(r"^[+-]?\d+\.\d+$", c)):
|
|
return float(c)
|
|
if c.lower() in ("false", ""):
|
|
return False
|
|
if c.lower() == "true":
|
|
return True
|
|
# Bare identifier - resolve as variable reference
|
|
varname, varspecifier = self.__separate_arrayref(c)
|
|
val = self.__get_variable_value(varname, varspecifier)
|
|
if val is None:
|
|
val = ""
|
|
vartype = "system" if self.state.variables.name_is_system(c) else "user"
|
|
self.warn_or_stop(f"Unknown {vartype} variable '{escape_single_quotes(c)}'")
|
|
if isinstance(val, str):
|
|
try:
|
|
return int(val)
|
|
except ValueError:
|
|
pass
|
|
if bool(re.match(r"^[+-]?\d+\.\d+$", val)):
|
|
return float(val)
|
|
if val in ("false", ""):
|
|
return False
|
|
if val == "true":
|
|
return True
|
|
val = val.lower()
|
|
return val
|
|
|
|
def __wmt(self, cond_desc, a, b, op):
|
|
return self.__warn_mixedtype(cond_desc, a, b, op)
|
|
|
|
def __pairwise_all(self, cond_desc, op1v, op2v, op):
|
|
return all(self.__wmt(cond_desc, a, b, op) for a, b in zip(op1v, op2v))
|
|
|
|
def __alltoone_all(self, cond_desc, op1v, op2v, op):
|
|
if isinstance(op1v, list):
|
|
return all(self.__wmt(cond_desc, a, op2v, op) for a in op1v)
|
|
else:
|
|
return all(self.__wmt(cond_desc, op1v, b, op) for b in op2v)
|
|
|
|
def __alltoone_any(self, cond_desc, op1v, op2v, op):
|
|
if isinstance(op1v, list):
|
|
return any(self.__wmt(cond_desc, a, op2v, op) for a in op1v)
|
|
else:
|
|
return any(self.__wmt(cond_desc, op1v, b, op) for b in op2v)
|
|
|
|
def __eval_basiccondition(
|
|
self,
|
|
cond_desc: str,
|
|
operand1: str | list[str],
|
|
operator: str,
|
|
operand2: str | list[str],
|
|
) -> bool:
|
|
"""
|
|
Evaluate a condition based on the given operands and operator.
|
|
|
|
Args:
|
|
cond_desc (str): The description of the condition (for logging).
|
|
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:
|
|
if not operand1_isarray and not operand2_isarray:
|
|
operations = {
|
|
"eq": lambda: self.__wmt(cond_desc, operand1_value, operand2_value, lambda x, y: x == y),
|
|
"ne": lambda: self.__wmt(cond_desc, operand1_value, operand2_value, lambda x, y: x != y),
|
|
"gt": lambda: self.__wmt(cond_desc, operand1_value, operand2_value, lambda x, y: x > y),
|
|
"lt": lambda: self.__wmt(cond_desc, operand1_value, operand2_value, lambda x, y: x < y),
|
|
"ge": lambda: self.__wmt(cond_desc, operand1_value, operand2_value, lambda x, y: x >= y),
|
|
"le": lambda: self.__wmt(cond_desc, operand1_value, operand2_value, lambda x, y: x <= y),
|
|
"in": lambda: self.__wmt(cond_desc, str(operand1_value), str(operand2_value), lambda x, y: x in y),
|
|
"any_in": None, # does not make sense
|
|
"contains": lambda: self.__wmt(
|
|
cond_desc, str(operand1_value), str(operand2_value), lambda x, y: y in x
|
|
),
|
|
"contains_any": None, # does not make sense
|
|
}
|
|
elif operand1_isarray and operand2_isarray:
|
|
operations = {
|
|
"eq": lambda: len(operand1_value) == len(operand2_value)
|
|
and self.__pairwise_all(cond_desc, operand1_value, operand2_value, lambda x, y: x == y),
|
|
"ne": lambda: len(operand1_value) != len(operand2_value)
|
|
or self.__pairwise_all(cond_desc, operand1_value, operand2_value, lambda x, y: x != y),
|
|
"gt": lambda: len(operand1_value) == len(operand2_value)
|
|
and self.__pairwise_all(cond_desc, operand1_value, operand2_value, lambda x, y: x > y),
|
|
"lt": lambda: len(operand1_value) == len(operand2_value)
|
|
and self.__pairwise_all(cond_desc, operand1_value, operand2_value, lambda x, y: x < y),
|
|
"ge": lambda: len(operand1_value) == len(operand2_value)
|
|
and self.__pairwise_all(cond_desc, operand1_value, operand2_value, lambda x, y: x >= y),
|
|
"le": lambda: len(operand1_value) == len(operand2_value)
|
|
and self.__pairwise_all(cond_desc, operand1_value, operand2_value, lambda x, y: x <= y),
|
|
"in": lambda: self.__alltoone_all(cond_desc, operand1_value, operand2_value, lambda x, y: x in y),
|
|
# all(
|
|
# self.__wmt(cond_desc, a, operand2_value, lambda x, y: x in y) for a in operand1_value
|
|
# ),
|
|
"any_in": lambda: self.__alltoone_any(
|
|
cond_desc, operand1_value, operand2_value, lambda x, y: x in y
|
|
),
|
|
# any(
|
|
# self.__wmt(cond_desc, a, operand2_value, lambda x, y: x in y) for a in operand1_value
|
|
# ),
|
|
"contains": lambda: self.__alltoone_all(
|
|
cond_desc, operand2_value, operand1_value, lambda x, y: x in y
|
|
),
|
|
# all(
|
|
# self.__wmt(cond_desc, a, operand1_value, lambda x, y: x in y) for a in operand2_value
|
|
# ),
|
|
"contains_any": lambda: self.__alltoone_any(
|
|
cond_desc, operand2_value, operand1_value, lambda x, y: x in y
|
|
),
|
|
# any(
|
|
# self.__wmt(cond_desc, a, operand1_value, lambda x, y: x in y) for a in operand2_value
|
|
# ),
|
|
}
|
|
elif operand1_isarray and not operand2_isarray:
|
|
if self.state.options.strict_operators:
|
|
operations = {
|
|
"eq": None, # does not make sense
|
|
"ne": None, # does not make sense
|
|
"gt": None, # does not make sense
|
|
"lt": None, # does not make sense
|
|
"ge": None, # does not make sense
|
|
"le": None, # does not make sense
|
|
}
|
|
else:
|
|
operations = {
|
|
"eq": lambda: self.__alltoone_all(
|
|
cond_desc, operand1_value, operand2_value, lambda x, y: x == y
|
|
),
|
|
"ne": lambda: self.__alltoone_all(
|
|
cond_desc, operand1_value, operand2_value, lambda x, y: x != y
|
|
),
|
|
"gt": lambda: self.__alltoone_all(
|
|
cond_desc, operand1_value, operand2_value, lambda x, y: x > y
|
|
),
|
|
"lt": lambda: self.__alltoone_all(
|
|
cond_desc, operand1_value, operand2_value, lambda x, y: x < y
|
|
),
|
|
"ge": lambda: self.__alltoone_all(
|
|
cond_desc, operand1_value, operand2_value, lambda x, y: x >= y
|
|
),
|
|
"le": lambda: self.__alltoone_all(
|
|
cond_desc, operand1_value, operand2_value, lambda x, y: x <= y
|
|
),
|
|
}
|
|
operations.update(
|
|
{
|
|
"in": lambda: self.__alltoone_all(
|
|
cond_desc, operand1_value, operand2_value, lambda x, y: str(x) in str(y)
|
|
),
|
|
"any_in": lambda: self.__alltoone_any(
|
|
cond_desc, operand1_value, operand2_value, lambda x, y: str(x) in str(y)
|
|
),
|
|
"contains": lambda: self.__wmt(cond_desc, operand2_value, operand1_value, lambda x, y: x in y),
|
|
"contains_any": None, # does not make sense
|
|
}
|
|
)
|
|
elif not operand1_isarray and operand2_isarray:
|
|
if self.state.options.strict_operators:
|
|
operations = {
|
|
"eq": None, # does not make sense
|
|
"ne": None, # does not make sense
|
|
"gt": None, # does not make sense
|
|
"lt": None, # does not make sense
|
|
"ge": None, # does not make sense
|
|
"le": None, # does not make sense
|
|
}
|
|
else:
|
|
operations = {
|
|
"eq": lambda: self.__alltoone_all(
|
|
cond_desc, operand1_value, operand2_value, lambda x, y: x == y
|
|
),
|
|
"ne": lambda: self.__alltoone_all(
|
|
cond_desc, operand1_value, operand2_value, lambda x, y: x != y
|
|
),
|
|
"gt": lambda: self.__alltoone_all(
|
|
cond_desc, operand1_value, operand2_value, lambda x, y: x > y
|
|
),
|
|
"lt": lambda: self.__alltoone_all(
|
|
cond_desc, operand1_value, operand2_value, lambda x, y: x < y
|
|
),
|
|
"ge": lambda: self.__alltoone_all(
|
|
cond_desc, operand1_value, operand2_value, lambda x, y: x >= y
|
|
),
|
|
"le": lambda: self.__alltoone_all(
|
|
cond_desc, operand1_value, operand2_value, lambda x, y: x <= y
|
|
),
|
|
}
|
|
operations.update(
|
|
{
|
|
"in": lambda: self.__wmt(cond_desc, operand1_value, operand2_value, lambda x, y: x in y),
|
|
"any_in": None, # does not make sense
|
|
"contains": lambda: self.__alltoone_all(
|
|
cond_desc, operand2_value, operand1_value, lambda x, y: str(x) in str(y)
|
|
),
|
|
"contains_any": lambda: self.__alltoone_any(
|
|
cond_desc, operand2_value, operand1_value, lambda x, y: str(x) in str(y)
|
|
),
|
|
}
|
|
)
|
|
else:
|
|
operations = {}
|
|
operation = operations.get(operator, None)
|
|
if operation is None:
|
|
self.warn_or_stop(
|
|
f"Unsupported operator '{escape_single_quotes(operator)}' in condition '{escape_single_quotes(cond_desc)}'"
|
|
)
|
|
return False
|
|
result = operation()
|
|
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_specifier = None
|
|
if vardescriptor.children[1] is not None:
|
|
vardescriptor_name += "[]"
|
|
if vardescriptor.children[2] is not None:
|
|
if isinstance(vardescriptor.children[2], lark.Token):
|
|
vardescriptor_specifier = str(vardescriptor.children[2])
|
|
else:
|
|
vardescriptor_specifier = "".join(vardescriptor.children[2].children)
|
|
return vardescriptor_name, vardescriptor_specifier
|
|
|
|
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_specifier = self.__separate_vardescriptor(value_node)
|
|
return (
|
|
vardescriptor_name[0:-2] + f"[{vardescriptor_specifier}]"
|
|
if vardescriptor_specifier 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 | 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"
|
|
cond_desc = (
|
|
cond_operand1
|
|
if isinstance(cond_operand1, str)
|
|
else "(" + ", ".join(str(c) for c in cond_operand1) + ")"
|
|
)
|
|
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_desc = f"{cond_operand1} {cond_operation} {cond_operand2 if isinstance(cond_operand2, str) else '(' + ', '.join(str(c) for c in cond_operand2) + ')'}"
|
|
cond_result = self.__eval_basiccondition(cond_desc, cond_operand1, cond_operation, cond_operand2)
|
|
if invert:
|
|
cond_result = not cond_result
|
|
return cond_result
|
|
|
|
def negative_sep(self, tree: lark.Tree):
|
|
"""
|
|
Process a negative prompt separator in the tree.
|
|
"""
|
|
start_result = self.__result
|
|
t1 = time.monotonic_ns()
|
|
x = tree.children[0]
|
|
self.__result += x.value
|
|
self.__is_negative = True
|
|
t2 = time.monotonic_ns()
|
|
self.__debug_end("negative_sep", start_result, t2 - t1)
|
|
|
|
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.and_
|
|
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.scheduling
|
|
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(TreeProcessor.AccumulatedShell("sc", pos))
|
|
self.__result += "["
|
|
if before is not None:
|
|
self.log(logging.DEBUG, f"Shell scheduled before with position {pos}")
|
|
self.__shell.append(TreeProcessor.AccumulatedShell("scb", pos))
|
|
self.__visit(before)
|
|
self.__shell.pop()
|
|
self.log(logging.DEBUG, f"Shell scheduled after with position {pos}")
|
|
self.__shell.append(TreeProcessor.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.alternation
|
|
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(TreeProcessor.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(TreeProcessor.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)
|
|
|
|
@staticmethod
|
|
def _try_extract_attention(s: str) -> tuple[str, float] | None:
|
|
"""
|
|
If s is entirely a single attention wrapper - (inner), (inner:W), or [inner] -
|
|
return (inner_content, weight). Otherwise return None.
|
|
|
|
Used for post-visit merging when a wildcard or choices construct expands to a
|
|
single attention that can be merged with an enclosing outer attention.
|
|
"""
|
|
if not s:
|
|
return None
|
|
open_char = s[0]
|
|
if open_char == "(":
|
|
close_char = ")"
|
|
elif open_char == "[":
|
|
close_char = "]"
|
|
else:
|
|
return None
|
|
depth = 0
|
|
for i, c in enumerate(s):
|
|
if c == open_char:
|
|
depth += 1
|
|
elif c == close_char:
|
|
depth -= 1
|
|
if depth == 0:
|
|
if i != len(s) - 1:
|
|
return None # wrapper closes before end of string -> multiple items
|
|
break
|
|
else:
|
|
return None # never fully closed
|
|
inner = s[1:-1]
|
|
if open_char == "[":
|
|
# Disambiguate from alternation [a|b] and scheduling [before:after:N].
|
|
# Both use [...] but are not attention constructs.
|
|
paren_depth = 0
|
|
bracket_depth = 0
|
|
top_level_pipes = 0
|
|
top_level_colons = 0
|
|
last_colon_pos = -1
|
|
for i, c in enumerate(inner):
|
|
if c == "(":
|
|
paren_depth += 1
|
|
elif c == ")":
|
|
paren_depth -= 1
|
|
elif c == "[":
|
|
bracket_depth += 1
|
|
elif c == "]":
|
|
bracket_depth -= 1
|
|
elif paren_depth == 0 and bracket_depth == 0:
|
|
if c == "|":
|
|
top_level_pipes += 1
|
|
elif c == ":":
|
|
top_level_colons += 1
|
|
last_colon_pos = i
|
|
if top_level_pipes > 0:
|
|
return None # alternation construct
|
|
if top_level_colons >= 2 and last_colon_pos >= 0:
|
|
try:
|
|
float(inner[last_colon_pos + 1 :])
|
|
return None # scheduling construct: [before:after:N]
|
|
except ValueError:
|
|
pass
|
|
return (inner, 0.9)
|
|
# Parenthesis form - scan backwards for a top-level :weight suffix
|
|
depth = 0
|
|
for i in range(len(inner) - 1, -1, -1):
|
|
c = inner[i]
|
|
if c in ")]":
|
|
depth += 1
|
|
elif c in "([":
|
|
depth -= 1
|
|
elif c == ":" and depth == 0:
|
|
try:
|
|
w = float(inner[i + 1 :])
|
|
return (inner[:i], w)
|
|
except ValueError:
|
|
break
|
|
return (inner, 1.1)
|
|
|
|
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:
|
|
# we check while the children are attentions, in which case we merge the weights
|
|
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
|
|
self.log(
|
|
logging.DEBUG,
|
|
f"Merging nested attention with weight {inner_weight}, cumulative weight now {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.attention
|
|
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(TreeProcessor.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})"
|
|
# Post-visit merge: if the entire visited content is a single attention wrapper
|
|
# (e.g. from a wildcard or choices expansion), merge weights here.
|
|
# The static tree-walk above only covers direct attention children; this
|
|
# handles the case where the inner attention came from an expanded wildcard.
|
|
if self.state.options.cup_merge_attention:
|
|
visited_content = self.__result[len(start_result) + len(starttag) :]
|
|
merge = TreeProcessor._try_extract_attention(visited_content)
|
|
if merge is not None:
|
|
inner_content, inner_weight = merge
|
|
weight = math.floor(weight * inner_weight * 100) / 100
|
|
self.log(
|
|
logging.DEBUG,
|
|
f"Merging nested attention with weight {inner_weight}, cumulative weight now {weight}",
|
|
)
|
|
weight_str = f"{weight:.2f}".rstrip("0").rstrip(".")
|
|
if weight_str == "1.1":
|
|
weight_kind = 2
|
|
starttag = "("
|
|
endtag = ")"
|
|
elif weight_str == "0.9" and attention_processing != "parentheses":
|
|
weight_kind = 1
|
|
starttag = "["
|
|
endtag = "]"
|
|
else:
|
|
weight_kind = 3
|
|
starttag = "("
|
|
endtag = f":{weight_str})"
|
|
self.__result = start_result + starttag + inner_content
|
|
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_nodes = tree.children[1::]
|
|
attention_processing = self.state.host_config.attention
|
|
peeled = False
|
|
if (
|
|
self.state.options.cup_merge_attention
|
|
and attention_processing in ("ok", "parentheses")
|
|
and len(content_nodes) == 1
|
|
and isinstance(content_nodes[0], lark.Tree)
|
|
and content_nodes[0].data == "attention"
|
|
):
|
|
inner_tree = content_nodes[0]
|
|
weight = 1.0
|
|
while isinstance(inner_tree, lark.Tree) and inner_tree.data == "attention":
|
|
if len(inner_tree.children) == 2:
|
|
w = inner_tree.children[-1]
|
|
inner_weight = float(w) if w is not None else 1.1
|
|
else:
|
|
inner_weight = 0.9
|
|
weight *= inner_weight
|
|
inner_tree = inner_tree.children[0]
|
|
weight = math.floor(weight * 100) / 100
|
|
weight_str = f"{weight:.2f}".rstrip("0").rstrip(".")
|
|
if weight_str == "0.9" and attention_processing != "parentheses":
|
|
weight_kind = 1
|
|
elif weight_str == "1.1":
|
|
weight_kind = 2
|
|
else:
|
|
weight_kind = 3
|
|
if attention_processing == "parentheses" and weight_kind == 1:
|
|
weight_kind = 3
|
|
weight_str = "0.9"
|
|
self.__shell.append(TreeProcessor.AccumulatedShell("at", (weight_kind, weight_str)))
|
|
content = self.__visit(inner_tree, False, True)
|
|
peeled = True
|
|
else:
|
|
content = self.__visit(content_nodes, False, True)
|
|
self.__negtags.append(
|
|
TreeProcessor.NegTag(len(self.__result), len(self.__result), content, parameters, self.__shell.copy())
|
|
)
|
|
if peeled:
|
|
self.__shell.pop()
|
|
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(
|
|
TreeProcessor.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_specifier: 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 self.state.variables.name_is_system(variable_name):
|
|
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_specifier 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_specifier}]"
|
|
value_description = self.__get_original_node_content(content, None)
|
|
value = content
|
|
raw_oldvalue = self.state.variables.get_user(variable_name)
|
|
newvalue = None
|
|
some_error = False
|
|
if variable_specifier is not None:
|
|
if (
|
|
is_array
|
|
and raw_oldvalue is not None
|
|
and isinstance(raw_oldvalue, list)
|
|
and not 0 <= int(variable_specifier) < len(raw_oldvalue)
|
|
):
|
|
self.warn_or_stop(
|
|
f"Invalid index {variable_specifier} 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_specifier 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_specifier = self.__separate_vardescriptor(
|
|
newvalue.children[0]
|
|
)
|
|
if vardescriptor_specifier is not None:
|
|
newvalue = None
|
|
else:
|
|
newvalue = self.__get_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 = ""
|
|
elif not isinstance(newvalue, list):
|
|
self.warn_or_stop(
|
|
f"Invalid starred value for '{escape_single_quotes(variable_name)}'! Starred values should be a list."
|
|
)
|
|
newvalue = ""
|
|
if is_array:
|
|
if variable_specifier is not None:
|
|
# Accessing an existing index, we need to update the array
|
|
newarray = raw_oldvalue.copy()
|
|
newarray[int(variable_specifier)] = 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.state.variables.set_user(variable_name, newvalue)
|
|
currentvalue = self.__get_variable_value(variable_name, variable_specifier, 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_specifier = self.__separate_vardescriptor(tree.children[0])
|
|
self.__varset("variableset", vardescriptor_name, vardescriptor_specifier, 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_specifier = self.__separate_vardescriptor(tree.children[0])
|
|
self.__varset("commandset", vardescriptor_name, vardescriptor_specifier, tree.children[1], tree.children[2])
|
|
|
|
def __varecho(
|
|
self,
|
|
command: str,
|
|
variable_name: str,
|
|
variable_specifier: 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_specifier}]" if variable_specifier is not None else variable_name
|
|
value = self.__get_variable_value(variable_name, variable_specifier, 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")
|
|
value = self.__visit(default, False, True)
|
|
self.__result += value
|
|
default_value = value
|
|
else:
|
|
self.warn_or_stop(f"Unknown variable {escape_single_quotes(vname)}")
|
|
default_value = ""
|
|
value = ""
|
|
if not self.state.variables.name_is_system(variable_name):
|
|
self.state.variables.echo(vname, value)
|
|
t2 = time.monotonic_ns()
|
|
info = variable_name
|
|
if is_array and variable_specifier is not None:
|
|
info = info[0:-2] + f"[{variable_specifier}]"
|
|
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_specifier = self.__separate_vardescriptor(tree.children[0])
|
|
self.__varecho(
|
|
"variableuse",
|
|
vardescriptor_name,
|
|
vardescriptor_specifier,
|
|
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_specifier = self.__separate_vardescriptor(tree.children[0])
|
|
self.__varecho(
|
|
"commandecho",
|
|
vardescriptor_name,
|
|
vardescriptor_specifier,
|
|
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:
|
|
cond = str(v.condition) if v.condition is not None else None
|
|
if cond:
|
|
try:
|
|
cnd = parse_prompt(
|
|
self.state,
|
|
"condition",
|
|
cond,
|
|
self.state.parsers["condition"],
|
|
True,
|
|
)
|
|
except lark.exceptions.UnexpectedInput as e:
|
|
self.warn_or_stop(
|
|
f"Error parsing condition '{escape_single_quotes(cond)}' 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 cond:
|
|
found_mappings.append(v)
|
|
else:
|
|
else_mapping = v
|
|
num_mappings = len(found_mappings)
|
|
if num_mappings > 0:
|
|
if self.state.options.do_combinatorial:
|
|
decision_idx = len(self.__comb_trace)
|
|
self.__comb_trace.append(num_mappings)
|
|
chosen_idx = (
|
|
min(self.__comb_forced_path[decision_idx], num_mappings - 1)
|
|
if decision_idx < len(self.__comb_forced_path)
|
|
else 0
|
|
)
|
|
found = found_mappings[chosen_idx]
|
|
elif num_mappings == 1:
|
|
found = found_mappings[0]
|
|
else:
|
|
found = found_mappings[
|
|
self.__rng.choice(
|
|
num_mappings,
|
|
p=[v.weight or 1 for v in found_mappings],
|
|
)
|
|
]
|
|
else:
|
|
found = else_mapping
|
|
if num_mappings < 2:
|
|
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[lark.Tree, list[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[lark.Tree,list[str]]: The resulting container and list of chosen choices.
|
|
"""
|
|
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 not in ("~", "@"):
|
|
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)
|
|
comb_chosen_selection: Optional[list[dict]] = None
|
|
if self.state.options.do_combinatorial or sampler == "@":
|
|
# Enumerate every distinct selection of choices, accounting for count range and repetition.
|
|
all_selections: list[tuple] = []
|
|
# When keep_choices_order is False the output depends on the selection order,
|
|
# so we must enumerate ordered sequences (permutations / product).
|
|
# When keep_choices_order is True selections are sorted afterward, so all
|
|
# orderings of the same items produce identical output and we only need
|
|
# unordered iterators (combinations / combinations_with_replacement).
|
|
for k in range(from_value, to_value + 1):
|
|
if repeating:
|
|
if self.state.options.keep_choices_order:
|
|
all_selections.extend(combinations_with_replacement(available_choices, k))
|
|
else:
|
|
all_selections.extend(product(available_choices, repeat=k))
|
|
else:
|
|
if self.state.options.keep_choices_order:
|
|
all_selections.extend(combinations(available_choices, k))
|
|
else:
|
|
all_selections.extend(permutations(available_choices, k))
|
|
num_selections = len(all_selections)
|
|
if self.state.options.do_combinatorial:
|
|
decision_idx = len(self.__comb_trace)
|
|
self.__comb_trace.append(num_selections)
|
|
chosen_idx = (
|
|
min(self.__comb_forced_path[decision_idx], num_selections - 1)
|
|
if decision_idx < len(self.__comb_forced_path)
|
|
else 0
|
|
)
|
|
else: # sampler == "@"
|
|
cycl_decision_idx = len(self.__cycl_trace)
|
|
self.__cycl_trace.append(num_selections)
|
|
chosen_idx = (
|
|
self.__cycl_forced_path[cycl_decision_idx] % num_selections
|
|
if cycl_decision_idx < len(self.__cycl_forced_path)
|
|
else 0
|
|
)
|
|
comb_chosen_selection = list(all_selections[chosen_idx])
|
|
num_choices = len(comb_chosen_selection)
|
|
else:
|
|
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:
|
|
if comb_chosen_selection is not None:
|
|
selected_choices: list[dict] = comb_chosen_selection
|
|
else:
|
|
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 = []
|
|
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)
|
|
# remove comments
|
|
results = [re.sub(r"\s*#[^\n]*(?:\n|$)", "", r, flags=re.DOTALL) for r in selected_choices_text]
|
|
else:
|
|
results = []
|
|
container = options.get("container", None)
|
|
if container is None:
|
|
separator = options.get("separator", self.state.options.choice_separator)
|
|
container = lark.Tree(
|
|
lark.Token("RULE", "content"),
|
|
[
|
|
lark.Tree(
|
|
lark.Token("RULE", "variableuse"),
|
|
[
|
|
lark.Tree(
|
|
lark.Token("RULE", "vardescriptor_get"),
|
|
[
|
|
lark.Token("__identifier", "_choices"),
|
|
lark.Token("__openbracket", "["),
|
|
lark.Tree(
|
|
lark.Token("RULE", "separator_descriptor"),
|
|
[
|
|
lark.Token("__separatorflag", "&"),
|
|
lark.Token("STRING", "'" + separator + "'"),
|
|
],
|
|
),
|
|
lark.Token("__closebracket", "]"),
|
|
],
|
|
),
|
|
None,
|
|
],
|
|
),
|
|
],
|
|
)
|
|
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 container, results
|
|
|
|
def __apply_container(self, container: lark.Tree, choices: list[str]) -> str:
|
|
# we save the choices variable in case there are nested choices
|
|
old_choices = self.state.variables.get_system("_choices[]", None)
|
|
self.state.variables.set_system("_choices[]", choices)
|
|
joined_results = self.__visit(container, False, True)
|
|
# we restore the old choices variable
|
|
self.state.variables.set_system("_choices[]", old_choices)
|
|
return joined_results
|
|
|
|
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]
|
|
container = options.get("container", None)
|
|
container_kind = "specified"
|
|
if container is None or not isinstance(container, str):
|
|
container_kind = "assembled"
|
|
has_separator = "separator" in options
|
|
has_prefix = "prefix" in options
|
|
has_suffix = "suffix" in options
|
|
if has_separator or has_prefix or has_suffix:
|
|
separator = options.get("separator", self.state.options.choice_separator)
|
|
container = "${_choices[&'" + separator + "']}"
|
|
prefix = options.get("prefix", None)
|
|
if prefix is not None and isinstance(prefix, str):
|
|
if prefix != "" and re.match(r"\w", prefix[-1]):
|
|
prefix += " "
|
|
container = prefix + container
|
|
suffix = options.get("suffix", None)
|
|
if suffix is not None and isinstance(suffix, str):
|
|
if suffix != "" and re.match(r"\w", suffix[0]):
|
|
suffix = " " + suffix
|
|
container += suffix
|
|
options.pop("separator", None)
|
|
options.pop("prefix", None)
|
|
options.pop("suffix", None)
|
|
options.pop("container", None)
|
|
if container is not None:
|
|
try:
|
|
options["container"] = parse_prompt(
|
|
self.state, "choicevalue", container, self.state.parsers["choicevalue"], True
|
|
)
|
|
except lark.exceptions.UnexpectedInput as e:
|
|
self.warn_or_stop(
|
|
f"Error parsing choice {container_kind} container '{escape_single_quotes(container)}' 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.__is_negative))
|
|
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_specifier = self.__separate_vardescriptor(var_object.children[0])
|
|
variablename = vardescriptor_name
|
|
# variablevalue = self.__visit(var_object.children[1], False, True)
|
|
variablebackup = self.state.variables.get_user(variablename)
|
|
# self.state.variables.delete_user(variablename)
|
|
# self.state.variables.set_user(variablename, variablevalue)
|
|
self.__varset("wildcard", variablename, vardescriptor_specifier, None, var_object.children[1])
|
|
choice_values_all = []
|
|
for wildcard in selected_wildcards:
|
|
if wildcard is None:
|
|
self.__detectedWildcards.append((wc, self.__is_negative))
|
|
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
|
|
container, chosen_choices = self.__get_choices_select(
|
|
applied_options, choice_values_all, filter_specifier, wildcard_key
|
|
)
|
|
if chosen_choices:
|
|
self.__result += self.__apply_container(container, chosen_choices)
|
|
if wildcard_key in self.__wildcard_filters:
|
|
del self.__wildcard_filters[wildcard_key]
|
|
if variablename is not None:
|
|
self.state.variables.delete_user(variablename)
|
|
if variablebackup is not None:
|
|
self.state.variables.set_user(variablename, variablebackup)
|
|
elif self.state.options.if_wildcards != IFWILDCARDS_CHOICES.remove:
|
|
self.__detectedWildcards.append((wc, self.__is_negative))
|
|
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]]:
|
|
if (
|
|
len(filters.children) == 1
|
|
and len(filters.children[0].children) == 1
|
|
and isinstance(filters.children[0].children[0], lark.Tree)
|
|
):
|
|
# special case when the whole filter is in a variable
|
|
# note that an individual label in a variable will also come through here
|
|
f = self.__visit(filters.children[0].children[0], False, True)
|
|
filters = parse_prompt(
|
|
self.state,
|
|
"filter specifier",
|
|
str(f),
|
|
self.state.parsers["wc_filter_or"],
|
|
True,
|
|
)
|
|
filter_specifier = []
|
|
for or_ in filters.children:
|
|
and_group = []
|
|
for and_ in or_.children:
|
|
label = and_.children[0]
|
|
if isinstance(label, lark.Token):
|
|
# it's a literal, we can use it directly
|
|
and_group.append(str(label))
|
|
else:
|
|
# it's a variable, we need to evaluate it
|
|
v = self.__visit(label, False, True)
|
|
# we remove commas and pluses to avoid confusion with the filter specifier syntax
|
|
and_group.append(v.replace(",", "").replace("+", ""))
|
|
filter_specifier.append(and_group)
|
|
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:")
|
|
container, chosen_choices = self.__get_choices_select(options, choice_values)
|
|
self.__result += self.__apply_container(container, chosen_choices)
|
|
elif self.state.options.if_wildcards != IFWILDCARDS_CHOICES.remove:
|
|
self.__detectedWildcards.append((ch, self.__is_negative))
|
|
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 __process_negtags(self):
|
|
# 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 self.state.host_config.attention != "parentheses":
|
|
new_kind = 1
|
|
elif new_weight_str == "1.1":
|
|
new_kind = 2
|
|
else:
|
|
new_kind = 3
|
|
negtag.shell[i - 1] = TreeProcessor.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}")
|
|
self.__negtags = []
|
|
|
|
def __apply_stn_insertions(self):
|
|
"""
|
|
Apply all accumulated STN content from add_at to self.result using the recorded
|
|
insertion_at positions, then reset both so ppp.py does not re-apply them.
|
|
"""
|
|
pos, neg = self.__result.split(self.NEGATIVE_SEP, 1)
|
|
neg_start = len(pos) + len(self.NEGATIVE_SEP)
|
|
stn_sep = self.state.options.stn_separator
|
|
self.log(logging.DEBUG, f"Applying STN additions to negative: {self.__add_at}")
|
|
self.log(logging.DEBUG, f"Applying STN indexes: {self.__insertion_at}")
|
|
ordered_range = sorted(
|
|
range(10),
|
|
key=lambda x: self.__insertion_at[x][0] if self.__insertion_at[x] is not None else float("-inf"),
|
|
reverse=True,
|
|
)
|
|
for n in ordered_range:
|
|
if self.__insertion_at[n] is not None:
|
|
insertion_point_n: list[str] = self.__add_at["insertion_point"][n]
|
|
ipp = self.__insertion_at[n][0] - neg_start
|
|
ipl = self.__insertion_at[n][1] - self.__insertion_at[n][0]
|
|
if neg[ipp - len(stn_sep) : ipp] == stn_sep:
|
|
ipp -= len(stn_sep) # adjust for existing start separator
|
|
ipl += len(stn_sep)
|
|
insertion_point_n.insert(0, neg[:ipp])
|
|
if neg[ipp + ipl : ipp + ipl + len(stn_sep)] == stn_sep:
|
|
ipl += len(stn_sep) # adjust for existing end separator
|
|
end_part = neg[ipp + ipl :]
|
|
if len(end_part) > 0:
|
|
insertion_point_n.append(end_part)
|
|
neg = stn_sep.join(insertion_point_n)
|
|
else:
|
|
ipp = 0
|
|
if neg.startswith(stn_sep):
|
|
ipp = len(stn_sep)
|
|
self.__add_at["insertion_point"][n].append(neg[ipp:])
|
|
neg = stn_sep.join(self.__add_at["insertion_point"][n])
|
|
if self.__add_at["start"]:
|
|
add_at_start = self.__add_at["start"]
|
|
if len(neg) > 0:
|
|
ipp = 0
|
|
if neg.startswith(stn_sep):
|
|
ipp = len(stn_sep) # adjust for existing end separator
|
|
add_at_start.append(neg[ipp:])
|
|
neg = stn_sep.join(add_at_start)
|
|
if self.__add_at["end"]:
|
|
add_at_end = self.__add_at["end"]
|
|
if len(neg) > 0:
|
|
ipl = len(neg)
|
|
if neg.endswith(stn_sep):
|
|
ipl -= len(stn_sep) # adjust for existing start separator
|
|
add_at_end.insert(0, neg[:ipl])
|
|
neg = stn_sep.join(add_at_end)
|
|
# self.add_at = {"start": [], "insertion_point": [[] for _ in range(10)], "end": []}
|
|
# self.insertion_at = [None for _ in range(10)]
|
|
self.__result = pos + self.NEGATIVE_SEP + neg
|
|
|
|
def start(self, tree):
|
|
self.__result = ""
|
|
t1 = time.monotonic_ns()
|
|
self.__visit(tree.children)
|
|
self.__process_negtags()
|
|
if self.__is_negative:
|
|
self.__apply_stn_insertions()
|
|
t2 = time.monotonic_ns()
|
|
self.__debug_end("start", "", t2 - t1)
|