From b971ea8d37f938518a239d60064cc6d7b9e02649 Mon Sep 17 00:00:00 2001 From: Antonio Cordero Balcazar Date: Sat, 16 May 2026 16:07:19 +0200 Subject: [PATCH] * Variables refactoring. * Scalar variables returned with their correct type. * Fix adding to indexed variables. --- .github/instructions/python.instructions.md | 4 + docs/SYNTAX.md | 6 +- ppp.py | 22 +-- ppp_common.py | 2 +- ppp_tree.py | 161 ++++++++++++-------- ppp_utils.py | 19 +++ ppp_variables.py | 123 +++++++++------ tests/tests_varcomms.py | 62 ++++++-- 8 files changed, 265 insertions(+), 134 deletions(-) diff --git a/.github/instructions/python.instructions.md b/.github/instructions/python.instructions.md index 205bf7d..351718e 100644 --- a/.github/instructions/python.instructions.md +++ b/.github/instructions/python.instructions.md @@ -41,3 +41,7 @@ All Python code must be compatible with Python 3.10. Do not use language feature ## Paths Use `pathlib.Path` for filesystem paths instead of `str` paths or `os.path`. + +## Comments + +Do not use emdashes (—) in comments. Use a single dash (-) or parentheses instead. Also avoid any other typographical punctuation that is not basic ASCII. diff --git a/docs/SYNTAX.md b/docs/SYNTAX.md index e408fb1..28eb58e 100644 --- a/docs/SYNTAX.md +++ b/docs/SYNTAX.md @@ -202,6 +202,8 @@ The wildcard identifier supports globbing. The filter does not allow the `^` or The prompt has access to some system variables that contain model information, options, and other things. There is also the possibility of defining user variables. +Variable values `true` and `false` are considered a boolean, and numeric content is an integer or float. + All these variables can be used to output content or behave differently based on their values. ### System variables @@ -237,7 +239,7 @@ The format is: `value` These are the available optional modifiers: * `evaluate`: the value of the variable is evaluated at this moment, instead of when it is used. -* `add`: the value is added to the current value of the variable. It does not force an immediate evaluation of the old nor the added value. +* `add`: the value is "added" to the current value of the variable (depending on their type). When possible, it does not force an immediate evaluation of the old or added values. * `ifundefined`: the value will only be set if the variable is undefined. The `add` and `ifundefined` modifiers are mutually exclusive and cannot be used together. @@ -364,7 +366,7 @@ The operands can be a variable (`variable`, `array[]`, `array[index]`), a quoted When an operand is or contains a variable, it is resolved to the variable's current value before the operation. -Variable values `true` and `false` are considered a boolean, and numeric content is an integer or float. Except in substring operations indicated below. String comparisons are case insensitive. +String comparisons are case insensitive, and substring comparisons (see below) are always considered as strings. The operation can be preceded by `not` for readability, instead of using it in the front. diff --git a/ppp.py b/ppp.py index 25de6a6..50e46c4 100644 --- a/ppp.py +++ b/ppp.py @@ -23,7 +23,7 @@ from ppp_classes import ( PPPState, PPPStateOptions, ) -from ppp_variables import VariableRepository +from ppp_variables import VariableRepository, VariableEntry from ppp_logging import DEBUG_LEVEL, log from ppp_tree import TreeProcessor from ppp_utils import escape_single_quotes @@ -845,10 +845,10 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in def __postprocess_result( self, - result: tuple[str, list[tuple[str, bool]], tuple[dict[str, str | None], dict[str, str | None]]], + result: tuple[str, list[tuple[str, bool]], dict[str, VariableEntry]], ) -> tuple[str, str, dict[str, str | None]]: variables = {} - unified_prompt, rem_wildcards, (_, echoed_variables_snapshot) = result + unified_prompt, rem_wildcards, variables_snapshot = result # Split the unified prompt back into prompt and negative prompt split_parts = unified_prompt.split("\x1d", 1) @@ -862,13 +862,16 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in self.log(logging.INFO, f"Result prompt: {prompt}") self.log(logging.INFO, f"Result negative_prompt: {negative_prompt}") try: - # Get and clean variables - var_keys = sorted(echoed_variables_snapshot.keys()) + # Get and clean variables - prefer the explicitly echoed value; fall back to the evaluated value. + var_keys = sorted(variables_snapshot.keys()) for k in var_keys: - ev = echoed_variables_snapshot.get(k) - variables[k] = self.__cleanup(ev, 0) if self.state.options.cup_cleanup_variables else ev - - self.log(logging.DEBUG, f"Result variables: {variables}") + entry = variables_snapshot[k] + ev = entry.last_echoed_value if entry.last_echoed_value is not None else entry.value + if ev is not None: + if isinstance(ev, str) and self.state.options.cup_cleanup_variables: + ev = self.__cleanup(ev, 0) + variables[k] = ev + self.log(logging.INFO, f"Result variables: {variables}") # Result checks warnings = [] @@ -962,7 +965,6 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in list: A list of tuples, each containing the processed prompt, negative prompt, and all variables. """ self.state.variables.clear_user() - self.state.variables.clear_echoed() # Parse both prompts processor = TreeProcessor(self.state, rng, on_model_info_update=self.__on_model_info_update) diff --git a/ppp_common.py b/ppp_common.py index fa30aaf..e063ccc 100644 --- a/ppp_common.py +++ b/ppp_common.py @@ -222,7 +222,7 @@ def get_model_class_from_filename(filename: str) -> str: return "" header = json.loads(header_bytes) - # Build a mock state dict — detection only needs key names and shapes + # Build a mock state dict - detection only needs key names and shapes class _ShapeProxy: def __init__(self, shape): self.shape = shape diff --git a/ppp_tree.py b/ppp_tree.py index 878bb64..131ba54 100644 --- a/ppp_tree.py +++ b/ppp_tree.py @@ -6,15 +6,16 @@ import math import re import textwrap import time -from typing import Any, Callable, Optional +from typing import Callable, Optional import lark import numpy as np from ppp_classes import IFWILDCARDS_CHOICES, SUPPORTED_APPS, PPPState from ppp_enmappings import PPPENMappingVariant from ppp_logging import DEBUG_LEVEL, log -from ppp_utils import escape_single_quotes +from ppp_utils import escape_single_quotes, repr_value from ppp_common import parse_prompt, warn_or_stop +from ppp_variables import ScalarValue, VariableEntry from ppp_wildcards import PPPWildcard @@ -87,7 +88,7 @@ class TreeProcessor(lark.visitors.Interpreter): def start_visit( self, parsed: lark.Tree, - ) -> list[tuple[str, list[tuple[str, bool]], tuple[dict[str, Any], dict[str, str]]]]: + ) -> list[tuple[str, list[tuple[str, bool]], dict[str, VariableEntry]]]: """ 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. @@ -96,8 +97,8 @@ class TreeProcessor(lark.visitors.Interpreter): 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 + list[tuple[str, list[tuple[str,bool]], dict[str, VariableEntry]]]: 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. """ @@ -111,17 +112,17 @@ class TreeProcessor(lark.visitors.Interpreter): self.__cycl_forced_path = list(self.state.cyclical_state.current_path) self.__cycl_trace = [] self.visit(parsed) - self.__finalize_echoed_variables() + self.__finalize_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())] + return [(self.__result, self.__detectedWildcards, self.state.variables.backup_user())] # 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() + initial_vars = self.state.variables.backup_user() results: list[tuple[str, list[tuple[str, bool]], tuple]] = [] limit = self.state.options.combinatorial_limit @@ -130,12 +131,10 @@ class TreeProcessor(lark.visitors.Interpreter): self.__comb_forced_path = list(forced_path) self.__comb_trace = [] self.__reset_run_state() - self.state.variables.restore_user_and_echoed(initial_vars) + self.state.variables.restore_user(initial_vars) self.visit(parsed) - self.__finalize_echoed_variables() - results.append( - (self.__result, self.__detectedWildcards.copy(), self.state.variables.backup_user_and_echoed()) - ) + self.__finalize_variables() + results.append((self.__result, self.__detectedWildcards.copy(), self.state.variables.backup_user())) 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}") @@ -171,16 +170,19 @@ class TreeProcessor(lark.visitors.Interpreter): 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 + def __finalize_variables(self): + """ + Ensure all variables have either an echoed value or their value evaluated as + a scalar at the end of processing, so they are included in the output snapshots. + """ + for k in self.state.variables.all_user: + if self.state.variables.get_echoed_value(k) is None and not isinstance( + self.state.variables.get_user(k), ScalarValue + ): 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 + name, specifier = self.__separate_arrayref(k) + value = self.get_final_scalar_variable(name, specifier) + self.state.variables.set_user(k, value) def __visit( self, @@ -209,7 +211,7 @@ class TreeProcessor(lark.visitors.Interpreter): 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() + backup_vars = self.state.variables.backup_user() if node is not None: if isinstance(node, list): for child in node: @@ -233,7 +235,7 @@ class TreeProcessor(lark.visitors.Interpreter): 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) + self.state.variables.restore_user(backup_vars) return added_result def __get_original_node_content(self, node: lark.Tree | lark.Token, default=None) -> str: @@ -267,14 +269,20 @@ class TreeProcessor(lark.visitors.Interpreter): return None, specifier[2:-1], False if not specifier.isdecimal(): # bare identifier: resolve as variable - specifier = self.get_final_variable(specifier) + name, spe = self.__separate_arrayref(specifier) + specifier = self.__value_to_str(self.__get_variable_value(name, spe, True, False)) 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 + self, + name: str, + specifier: str | None = None, + evaluate=True, + add_to_content=False, + restore_state=True, ) -> str | int | float | bool | list | None: """ Get the value of a variable. @@ -283,7 +291,8 @@ class TreeProcessor(lark.visitors.Interpreter): 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). + add_to_content (bool): Whether to add the content after visiting. + restore_state (bool): Whether to restore the state after visiting. Returns: str|int|float|bool|list|None: The value of the variable, with strings coerced to their most specific type. @@ -293,15 +302,15 @@ class TreeProcessor(lark.visitors.Interpreter): visited = False if isinstance(v, lark.Tree): if evaluate: - v = self.__visit(v, restore_state=not visit, discard_content=not visit) - visited = visit + v = self.__visit(v, restore_state=restore_state, discard_content=not add_to_content) + visited = add_to_content else: - v = self.__get_original_node_content(v, "") + v = self.__get_original_node_content(v, "(cannot evaluate)") elif isinstance(v, lark.Token): v = str(v) if isinstance(v, str): v = self.__coerce_value(v) - if visit and not visited: + if add_to_content and not visited: self.__result += self.__value_to_str(v) return v @@ -314,7 +323,7 @@ class TreeProcessor(lark.visitors.Interpreter): idx, sep, cnt = self.__parse_array_specifier(specifier) if cnt: v = len(v) - if visit: + if add_to_content: self.__result += str(v) elif idx is not None: if 0 <= idx < len(v): @@ -327,10 +336,10 @@ class TreeProcessor(lark.visitors.Interpreter): v2 = [] for i, item in enumerate(v): v2.append(visit_value(item)) - if visit and i < len(v) - 1: + if add_to_content and i < len(v) - 1: self.__result += sep v = v2 - else: + elif not isinstance(v, str): # if it is a string we assume it has already been evaluated and joined v = None # error elif specifier is not None: v = None # error @@ -355,24 +364,24 @@ class TreeProcessor(lark.visitors.Interpreter): specifier = None return name, specifier - def get_final_variable(self, name_specifier: str) -> str: + def get_final_scalar_variable(self, name: str, specifier: str | None) -> ScalarValue: """ - Get the final value of a variable, resolving any references if needed. + Get the final scalar value of a variable, resolving any references if needed. Args: - name_specifier (str): The variable reference string. + name (str): The name of the variable. + specifier (str|None): The specifier for an array variable. Returns: - str: The final value of the variable. + ScalarValue: The final value of the variable. """ - name, specifier = self.__separate_arrayref(name_specifier) - v = self.__get_variable_value(name, specifier, True, False) + v = self.__get_variable_value(name, specifier, True, False, False) if isinstance(v, list): _, sep, _ = self.__parse_array_specifier(specifier) if sep is None: sep = self.state.options.choice_separator v = sep.join(self.__value_to_str(item) for item in v) - return self.__value_to_str(v) + return v def __debug_end(self, construct: str, start_result: str, duration: int, info=None): """ @@ -475,7 +484,7 @@ class TreeProcessor(lark.visitors.Interpreter): return True # Bare identifier - resolve as variable reference varname, varspecifier = self.__separate_arrayref(c) - val = self.__get_variable_value(varname, varspecifier) + val = self.__get_variable_value(varname, varspecifier, True, False) if val is None: val = "" vartype = "system" if self.state.variables.name_is_system(c) else "user" @@ -795,6 +804,7 @@ class TreeProcessor(lark.visitors.Interpreter): x = tree.children[0] self.__result += x.value self.__is_negative = True + self.__finalize_variables() # finalize variables here to capture their state at the point of separation t2 = time.monotonic_ns() self.__debug_end("negative_sep", start_result, t2 - t1) @@ -1269,6 +1279,10 @@ class TreeProcessor(lark.visitors.Interpreter): return info = f"{variable_name[0:-2]}[{variable_specifier}]" value_description = self.__get_original_node_content(content, None) + if value_description is None: + value_description = "" + else: + value_description = self.__coerce_value(value_description) value = content raw_oldvalue = self.state.variables.get_user(variable_name) newvalue = None @@ -1299,7 +1313,32 @@ class TreeProcessor(lark.visitors.Interpreter): 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 '')}'" + info += f" += {repr_value(value_description)}" + + def build_addition(old, added) -> lark.Tree: + if old is None: + return added + elif isinstance(old, ScalarValue): + if isinstance(added, ScalarValue): + evaluated_added = added + else: + evaluated_added = self.__coerce_value(self.__visit(added, False, True)) + if isinstance(evaluated_added, (str, int, float)): + if old.__class__ == evaluated_added.__class__: + return old + evaluated_added + return str(old) + str(evaluated_added) + return lark.Tree( + lark.Token("RULE", "varvalue"), + [lark.Token("plain", self.__value_to_str(old)), added], + # Meta should be {"content": str(old) + added.meta.content}, + ) + else: # lark.Tree + return lark.Tree( + lark.Token("RULE", "varvalue"), + [old, added], + # Meta should be {"content": old.meta.content + added.meta.content}, + ) + if raw_oldvalue is None: newvalue = value self.warn_or_stop(f"Unknown variable {escape_single_quotes(variable_name)}") @@ -1308,22 +1347,14 @@ class TreeProcessor(lark.visitors.Interpreter): self.warn_or_stop( f"Invalid variable value for '{escape_single_quotes(variable_name)}'! Cannot add to a non-array value." ) + elif variable_specifier is not None: + newvalue = build_addition(raw_oldvalue[int(variable_specifier)], 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}, - ) + else: + newvalue = build_addition(raw_oldvalue, value) elif any(item in modifiers_str for item in ["?", "ifundefined"]): - info += f" ?= '{escape_single_quotes(value_description or '')}'" + info += f" ?= {repr_value(value_description)}" if raw_oldvalue is None: newvalue = value else: @@ -1347,7 +1378,7 @@ class TreeProcessor(lark.visitors.Interpreter): if vardescriptor_specifier is not None: newvalue = None else: - newvalue = self.__get_variable_value(vardescriptor_name, None) + newvalue = self.__get_variable_value(vardescriptor_name, None, True, False) elif newvalue.children[0].data == "listvalue": newvalue = list( self.__resolve_operand(c) for c in self.__get_cond_operand(newvalue.children[0]) @@ -1389,13 +1420,14 @@ class TreeProcessor(lark.visitors.Interpreter): else: newvalue = raw_oldvalue + [newvalue] self.state.variables.set_user(variable_name, newvalue) - currentvalue = self.__get_variable_value(variable_name, variable_specifier, False) + currentvalue = self.__get_variable_value(variable_name, variable_specifier, False, False) if currentvalue is None: info += "error" elif isinstance(currentvalue, list): - info += "[" + ", ".join(f"'{escape_single_quotes(self.__value_to_str(v))}'" for v in currentvalue) + "]" + info_elems = [repr_value(v) for v in currentvalue] + info += "[" + ", ".join(info_elems) + "]" else: - info += f"'{escape_single_quotes(self.__value_to_str(currentvalue))}'" + info += repr_value(currentvalue) t2 = time.monotonic_ns() self.__debug_end(command, start_result, t2 - t1, info) @@ -1431,11 +1463,10 @@ class TreeProcessor(lark.visitors.Interpreter): 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) + # value = self.__get_variable_value(variable_name, variable_specifier, True, True) + value = self.get_final_scalar_variable(variable_name, variable_specifier) if value is None: if default is not None: self.log(logging.DEBUG, f"Variable '{escape_single_quotes(vname)}' not found, using default value") @@ -1446,8 +1477,10 @@ class TreeProcessor(lark.visitors.Interpreter): self.warn_or_stop(f"Unknown variable {escape_single_quotes(vname)}") default_value = "" value = "" + else: + self.__result += self.__value_to_str(value) if not self.state.variables.name_is_system(variable_name): - self.state.variables.echo(vname, value) + self.state.variables.set_echoed_value(vname, value) t2 = time.monotonic_ns() info = variable_name if is_array and variable_specifier is not None: diff --git a/ppp_utils.py b/ppp_utils.py index c032eb0..f9c1030 100644 --- a/ppp_utils.py +++ b/ppp_utils.py @@ -1,3 +1,6 @@ +from typing import Any + + def deep_freeze(obj): """ Deep freeze an object. @@ -43,6 +46,22 @@ def escape_double_quotes(s: str): return s.replace('"', '\\"') +def repr_value(s: Any): + """ + Return a string representation of a value, escaping single quotes. + + Args: + s (Any): The value to represent. + Returns: + str: The string representation of the value. + """ + if isinstance(s, str): + return f"'{escape_single_quotes(s)}'" + if isinstance(s, bool): + return "true" if s else "false" + return str(s) + + def format_output(text: str) -> str: """ Formats the output text by encoding it using unicode_escape and decoding it using utf-8. diff --git a/ppp_variables.py b/ppp_variables.py index cf28a5d..b47916e 100644 --- a/ppp_variables.py +++ b/ppp_variables.py @@ -1,23 +1,33 @@ +from dataclasses import dataclass, field from typing import Any +ScalarValue = str | int | float | bool +VariableValue = ScalarValue | list | None + + +@dataclass +class VariableEntry: + """Holds all state for a single user variable.""" + + value: Any = field(default=None) # raw unevaluated value or evaluated on set + last_echoed_value: ScalarValue | None = field(default=None) + class VariableRepository: """ - Unified repository for system, user, and echoed prompt variables. + Unified repository for system and user prompt variables. System variables (underscore-prefixed names like ``_model``) are populated once per processing session and are read-only during prompt evaluation. User variables are created and mutated by set/echo constructs in the prompt. - - Echoed variables record which user variables have already been output and - with what resolved string value. + Each user variable is stored as a :class:`VariableEntry` that tracks the + set value and the last value that was echoed into the prompt output. """ def __init__(self) -> None: - self._system: dict[str, Any] = {} - self._user: dict[str, Any] = {} - self._echoed: dict[str, str] = {} + self._system: dict[str, VariableValue] = {} + self._vars: dict[str, VariableEntry] = {} def name_is_system(self, name: str) -> bool: """Return True if *name* is a system variable (i.e. starts with an underscore).""" @@ -25,17 +35,17 @@ class VariableRepository: # ---- System variables ---- - def get_system(self, name: str, default: Any = None) -> Any: + def get_system(self, name: str, default: VariableValue = None) -> VariableValue: """Return the value of a system variable, or *default* if absent.""" return self._system.get(name, default) - def set_system(self, name: str, value: Any) -> None: + def set_system(self, name: str, value: VariableValue) -> None: """Set a system variable.""" if not self.name_is_system(name): raise ValueError(f"Invalid system variable name '{name}': must start with an underscore") self._system[name] = value - def update_system(self, mapping: dict[str, Any]) -> None: + def update_system(self, mapping: dict[str, VariableValue]) -> None: """Bulk-update system variables from *mapping*.""" for name in mapping: if not self.name_is_system(name): @@ -47,43 +57,79 @@ class VariableRepository: self._system.clear() @property - def all_system(self) -> dict[str, Any]: + def all_system(self) -> dict[str, VariableValue]: """Return a shallow copy of all system variables.""" return self._system.copy() # ---- User variables ---- + def _entry(self, name: str) -> VariableEntry: + """Return (creating if necessary) the :class:`VariableEntry` for *name*.""" + if name not in self._vars: + self._vars[name] = VariableEntry() + return self._vars[name] + def get_user(self, name: str, default: Any = None) -> Any: """Return the value of a user variable, or *default* if absent.""" - return self._user.get(name, default) + entry = self._vars.get(name) + if entry is None or entry.value is None: + return default + return entry.value def set_user(self, name: str, value: Any) -> None: - """Set a user variable.""" + """Set the value of a user variable.""" if self.name_is_system(name): raise ValueError(f"Invalid user variable name '{name}': must not start with an underscore") - self._user[name] = value + entry = self._entry(name) + entry.value = value def delete_user(self, name: str) -> None: - """Remove a user variable (no-op if it does not exist).""" - self._user.pop(name, None) + """ + Remove the value for a user variable. + """ + entry = self._vars.get(name) + if entry is None: + return + del self._vars[name] def clear_user(self) -> None: - """Remove all user variables.""" - self._user.clear() + """ + Clear the values for all user variables. + """ + self._vars.clear() - # ---- Echoed variables ---- + @property + def all_user(self) -> set[str]: + """Return the set of all user-variable keys (those with any non-None field).""" + return set(self._vars) - def get_echoed_value(self, name: str, default: str | None = None) -> str | None: - """Return the echoed string value for *name*, or *default* if not echoed.""" - return self._echoed.get(name, default) + def set_echoed_value(self, name: str, value: ScalarValue) -> None: + """Record that *name* was echoed into the prompt with *value*.""" + self._entry(name).last_echoed_value = value - def echo(self, name: str, value: str) -> None: - """Record that *name* was echoed with *value*.""" - self._echoed[name] = value + def get_echoed_value(self, name: str, default: ScalarValue | None = None) -> ScalarValue | None: + """Return the last echoed value for *name*, or *default* if it has not been echoed.""" + entry = self._vars.get(name) + if entry is None: + return default + return entry.last_echoed_value if entry.last_echoed_value is not None else default - def clear_echoed(self) -> None: - """Remove all echoed-variable records.""" - self._echoed.clear() + def backup_user(self) -> dict[str, VariableEntry]: + """Return a per-entry shallow-copy snapshot of all user variables for rollback.""" + return { + name: VariableEntry(entry.value, entry.last_echoed_value) + for name, entry in self._vars.items() + } + + def restore_user(self, backup: dict[str, VariableEntry]) -> None: + """Restore user variables from a snapshot made by :meth:`backup_user_and_echoed`.""" + self._vars.clear() + self._vars.update( + { + name: VariableEntry(entry.value, entry.last_echoed_value) + for name, entry in backup.items() + } + ) # ---- Combined queries ---- @@ -97,23 +143,4 @@ class VariableRepository: """ if name in self._system: return self._system.get(name, default) - return self._user.get(name, default) - - @property - def all_user_or_echoed_keys(self) -> set[str]: - """Return the union of user-variable and echoed-variable keys.""" - return set(self._user.keys()) | set(self._echoed.keys()) - - # ---- State backup / restore ---- - - def backup_user_and_echoed(self) -> tuple[dict[str, Any], dict[str, str]]: - """Return shallow-copy snapshots of user and echoed variables for rollback.""" - return self._user.copy(), self._echoed.copy() - - def restore_user_and_echoed(self, backup: tuple[dict[str, Any], dict[str, str]]) -> None: - """Restore user and echoed variables from a snapshot made by :meth:`backup_user_and_echoed`.""" - user_backup, echoed_backup = backup - self._user.clear() - self._user.update(user_backup) - self._echoed.clear() - self._echoed.update(echoed_backup) + return self.get_user(name, default) diff --git a/tests/tests_varcomms.py b/tests/tests_varcomms.py index bcac8b7..5404e59 100644 --- a/tests/tests_varcomms.py +++ b/tests/tests_varcomms.py @@ -21,7 +21,21 @@ class TestVarCommands(TestPromptPostProcessorBase): "${v1=}${v3:}", "", ), - OutputTuple("", "",{"v1": "", "v2": "", "v3": ""}), + OutputTuple("", "", {"v1": "", "v2": "", "v3": ""}), + ) + + # Typed and output variables test + def test_typed_output_variables(self): + self.process( + InputTuple( + "${str=value}${int=42}${float=3.14}${array[]=*('a','b','c')}${bool=true}${str2:default1},${str2:default2}", + "", + ), + OutputTuple( + "default1,default2", + "", + {"str": "value", "int": 42, "float": 3.14, "array[]": "a, b, c", "bool": True, "str2": "default2"}, + ), ) # Echoed variables tests @@ -99,31 +113,53 @@ class TestVarCommands(TestPromptPostProcessorBase): # Array variable tests - def test_array_variable_1(self): # array variable set with += and test of index value and full array with and without default separator + def test_array_variable_1( + self, + ): # array variable set with += and test of index value and full array with and without default separator self.process( InputTuple( "${v1[]=val1}${v1[]+=val2}${v1[]+=val3}${v1[1]:defval},${v1[]:defval2},${v1[&'.']:defval3}", "", ), - OutputTuple("val2,val1, val2, val3,val1.val2.val3", "", {"v1[]": "val1, val2, val3", "v1[1]": "val2", "v1[&'.']": "val1.val2.val3"}), + OutputTuple( + "val2,val1, val2, val3,val1.val2.val3", + "", + {"v1[]": "val1, val2, val3", "v1[1]": "val2", "v1[&'.']": "val1.val2.val3"}, + ), ) - def test_array_variable_2(self): # override of array variable value, test of default value when array variable is empty, test of default value when array variable is not set + def test_array_variable_2( + self, + ): # override of array variable value, test of default value when array variable is empty, test of default value when array variable is not set self.process( InputTuple( "${v1[]=val1}${v1[]=val2}${v1[]:defval},${v2[]:defval2},${v2[1]:defval3},${v3[]=}${v3[]:defval4}", "", ), - OutputTuple("val2,defval2,defval3", "", {"v1[]": "val2", "v2[]": "defval2", "v2[1]": "defval3", "v3[]": ""}), + OutputTuple( + "val2,defval2,defval3", "", {"v1[]": "val2", "v2[]": "defval2", "v2[1]": "defval3", "v3[]": ""} + ), ) - def test_array_variable_3(self): # access array index by variable, set array variable to expanded array variable and add expanded array + def test_array_variable_3( + self, + ): # access array index by variable, set array variable to expanded array variable and add expanded array self.process( InputTuple( "${v1[]=val1}${v1[]+=val2}${v2=1}${v1[v2]:defval1}${v3[]=${v1[]}}${v3[]+=${v1[]}}, ${v3[&'.']}", "", ), - OutputTuple("val2, val1, val2.val1, val2", "", {"v1[]": "val1, val2", "v2": "1", "v1[v2]": "val2", "v3[]": "val1, val2, val1, val2", "v3[&'.']": "val1, val2.val1, val2"}), + OutputTuple( + "val2, val1, val2.val1, val2", + "", + { + "v1[]": "val1, val2", + "v2": 1, + "v1[v2]": "val2", + "v3[]": "val1, val2, val1, val2", + "v3[&'.']": "val1, val2.val1, val2", + }, + ), ) def test_array_variable_4(self): # test list in array @@ -177,7 +213,7 @@ class TestVarCommands(TestPromptPostProcessorBase): "${v1[]=val1}${v1[]+=val2}${v1[]+=val3}${v1[#]:defval}, OKnot OK", "", ), - OutputTuple("3, OK", "", {"v1[]": "val1, val2, val3", "v1[#]": "3"}), + OutputTuple("3, OK", "", {"v1[]": "val1, val2, val3", "v1[#]": 3}), ) def test_array_variable_10(self): # array variable set with expanded values from wildcards in command format @@ -189,6 +225,15 @@ class TestVarCommands(TestPromptPostProcessorBase): OutputTuple("choice3", "", {"v1[]": "choice2, choice1, choice3, choice1"}), ) + def test_array_variable_11(self): # array variable set and indexed value added + self.process( + InputTuple( + "${v1[]=*(1,2,3)}${v1[0]+=10}${v2[]=*('1','2','3')}${v2[1]+=10}", + "", + ), + OutputTuple("", "", {"v1[]": "11, 2, 3", "v2[]": "1, 210, 3"}), + ) + # Operator tests ## R vs R @@ -636,7 +681,6 @@ class TestVarCommands(TestPromptPostProcessorBase): OutputTuple("OK", ""), ) - # NaN/undefined variable integer comparison tests def test_cmd_if_undefined_var_int_compare_warn(self): # undefined var integer compare with on_warning=warn