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

568 lines
23 KiB
Python

import fnmatch
from pathlib import Path
from typing import Any, Optional
import logging
from ruamel.yaml import YAML as _YAML
from ruamel.yaml.error import YAMLError as _YAMLError
from ppp_logging import DEBUG_LEVEL, log
from ppp_utils import deep_freeze, escape_single_quotes
class PPPWildcard:
"""
A wildcard object.
Attributes:
key (str): The key of the wildcard.
file (Path | None): The path to the file where the wildcard is defined, or None if from inline input.
unprocessed_choices (list[str]): The unprocessed choices of the wildcard.
options (dict): The options of the wildcard.
choices (list[dict]): The processed choices of the wildcard.
"""
def __init__(self, fullpath: Path | None, key: str, choices: list[str]):
self.key: str = key
self.file: Path | None = fullpath
self.unprocessed_choices: list[str] = choices
self.choices: list[dict] = None
self.options: dict = None
def __hash__(self) -> int:
t = (self.key, deep_freeze(self.unprocessed_choices))
return hash(t)
def __sizeof__(self):
return (
self.key.__sizeof__()
+ (self.file.__sizeof__() if self.file is not None else 0)
+ self.unprocessed_choices.__sizeof__()
+ self.choices.__sizeof__()
+ self.options.__sizeof__()
)
class PPPWildcards:
"""
A class to manage wildcards.
Attributes:
wildcards (dict[str, PPPWildcard]): The wildcards.
"""
DEFAULT_WILDCARDS_FOLDER = "wildcards"
def __init__(self, logger=None):
self.__logger: logging.Logger = logger
self.__debug_level = DEBUG_LEVEL.none
self.__wildcards_folders: list[Path] = []
self.__wildcard_files: dict[Path, float] = {}
self.__local_input_hash: int | None = None
self.__wildcard_default_filters: dict[str, list[list[str]]] = {}
self.wildcards: dict[str, PPPWildcard] = {}
def __hash__(self) -> int:
return hash(deep_freeze(self.wildcards))
def __sizeof__(self):
return self.wildcards.__sizeof__() + self.__wildcards_folders.__sizeof__() + self.__wildcard_files.__sizeof__()
def refresh_wildcards(
self,
debug_level: DEBUG_LEVEL,
wildcards_folders: Optional[list[Path]],
wildcards_input: str = None,
):
"""
Initialize the wildcards.
"""
self.reset_default_filters()
self.__debug_level = debug_level
self.__wildcards_folders = [Path(f) for f in (wildcards_folders or [])]
# log(self.__logger, self.__debug_level, logging.INFO, "Refreshing wildcards...")
# t1 = time.monotonic_ns()
for fullpath in list(self.__wildcard_files.keys()):
if not fullpath.exists() or not any(
fullpath.parent.is_relative_to(folder) for folder in self.__wildcards_folders
):
self.__remove_wildcards_from_path(fullpath)
if wildcards_input is None and self.__local_input_hash is not None:
self.__remove_wildcards_from_input()
if wildcards_folders is not None or wildcards_input is not None:
if wildcards_folders is not None:
for f in self.__wildcards_folders:
self.__get_wildcards_in_path(f if f.is_dir() else f.parent, f)
if wildcards_input is not None:
self.__get_wildcards_in_input(wildcards_input)
else:
self.wildcards = {}
self.__wildcard_files = {}
self.__local_input_hash = None
# t2 = time.monotonic_ns()
# log(self.__logger, self.__debug_level, logging.INFO, f"Wildcards refresh time: {(t2 - t1) / 1_000_000_000:.3f} seconds")
def get_wildcards(self, key: str) -> list[PPPWildcard]:
"""
Get all wildcards that match a key.
Args:
key (str): The key to match.
Returns:
list: A list of all wildcards that match the key.
"""
keys = sorted(fnmatch.filter(self.wildcards.keys(), key))
return [self.wildcards[k] for k in keys]
def __get_wc_in_dict(self, dictionary: dict, prefix="") -> list[tuple[str, Any]]:
"""
Get all wildcards in a dictionary, along their object.
Args:
dictionary (dict): The dictionary to check.
prefix (str): The prefix for the current key.
Returns:
list: A list of all leaf wildcards in the dictionary.
"""
wc = []
for key, obj in dictionary.items():
if isinstance(obj, dict):
wc.extend(self.__get_wc_in_dict(obj, prefix + str(key) + "/"))
else:
wc.append((prefix + str(key), obj))
return wc
def __remove_wildcards_from_path(self, full_path: Path, debug=True):
"""
Clear all wildcards from a file.
Args:
full_path (Path): The path to the file.
debug (bool): Whether to print debug messages or not.
"""
if debug and full_path in self.__wildcard_files:
log(
self.__logger,
self.__debug_level,
logging.DEBUG,
f"Removing from memory wildcards from file: {full_path}",
)
if full_path in self.__wildcard_files:
del self.__wildcard_files[full_path]
for key in list(self.wildcards.keys()):
if self.wildcards[key].file == full_path:
del self.wildcards[key]
def __remove_wildcards_from_input(self, debug=True):
"""
Clear all wildcards loaded from inline input.
Args:
debug (bool): Whether to print debug messages or not.
"""
if debug and self.__local_input_hash is not None:
log(self.__logger, self.__debug_level, logging.DEBUG, "Removing from memory wildcards from input")
self.__local_input_hash = None
for key in list(self.wildcards.keys()):
if self.wildcards[key].file is None:
del self.wildcards[key]
def __get_wildcards_in_file(self, base: Path, full_path: Path):
"""
Get all wildcards in a file.
Args:
base (Path): The base path for the wildcards.
full_path (Path): The path to the file.
"""
try:
last_modified = full_path.stat().st_mtime
last_modified_cached = self.__wildcard_files.get(full_path, None)
if last_modified_cached is not None and last_modified == self.__wildcard_files[full_path]:
return
extension = full_path.suffix
if extension not in (".txt", ".json", ".yaml", ".yml"):
return
self.__remove_wildcards_from_path(full_path, False)
if last_modified_cached is not None:
log(self.__logger, self.__debug_level, logging.DEBUG, f"Updating wildcards from file: {full_path}")
if extension == ".txt":
self.__get_wildcards_in_text_file(full_path, base)
elif extension in (".json", ".yaml", ".yml"):
self.__get_wildcards_in_structured_file(full_path, base)
self.__wildcard_files[full_path] = last_modified
except Exception as e: # pylint: disable=broad-except
log(
self.__logger,
self.__debug_level,
logging.ERROR,
f"Error reading wildcard file '{escape_single_quotes(str(full_path))}': {e}",
)
def __get_wildcards_in_input(self, wildcards_input: str):
"""
Get all wildcards in the string.
Args:
wildcards_input (str): The input string containing wildcards in json or yaml format.
"""
try:
new_h = hash(wildcards_input)
if new_h == self.__local_input_hash:
return
was_loaded = self.__local_input_hash is not None
self.__remove_wildcards_from_input(False)
if was_loaded:
log(self.__logger, self.__debug_level, logging.DEBUG, "Updating wildcards from input")
wildcards_input = wildcards_input.strip()
if wildcards_input != "":
try:
content = _YAML(typ='safe').load(wildcards_input)
except _YAMLError as e:
log(self.__logger, self.__debug_level, logging.WARNING, f"Invalid format for input wildcards: {e}")
return
if content is not None:
self.__add_wildcard(content, None, [""])
self.__local_input_hash = new_h
except Exception as e: # pylint: disable=broad-except
log(self.__logger, self.__debug_level, logging.ERROR, f"Error reading wildcards input: {e}")
# NOTE wcdef and choice options should not have properties in common
def is_dict_wcdef_options(self, d: dict) -> bool:
"""
Check if a dictionary is a valid wildcard definition options dictionary.
Args:
d (dict): The dictionary to check.
Returns:
bool: Whether the dictionary is a valid wildcard definition options dictionary or not.
"""
return all(
k
in [
"sampler",
"repeating",
"optional",
"count",
"from",
"to",
"prefix",
"suffix",
"container",
"description",
"separator",
]
for k in d.keys()
)
def is_dict_choice_options(self, d: dict) -> bool:
"""
Check if a dictionary is a valid choice options dictionary.
Args:
d (dict): The dictionary to check.
Returns:
bool: Whether the dictionary is a valid choice options dictionary or not.
"""
return all(k in ["command", "labels", "weight", "if", "content", "text"] for k in d.keys())
def __get_choices(self, obj: object, full_path: Path | None, key_parts: list[str]) -> list:
"""
We process the choices in the object and return them as a list.
Args:
obj (object): the value of a wildcard
full_path (Path | None): path to the file where the wildcard is defined, or None if from inline input
key_parts (list[str]): parts of the key for the wildcard
Returns:
list: list of choices
"""
if obj is None:
return None
if isinstance(obj, (str, dict)):
return [obj]
if isinstance(obj, (int, float, bool)):
return [str(obj)]
file_str = str(full_path) if full_path is not None else "input"
if not isinstance(obj, list) or len(obj) == 0:
log(
self.__logger,
self.__debug_level,
logging.WARNING,
f"Invalid format in wildcard '{escape_single_quotes('/'.join(key_parts))}' in file '{escape_single_quotes(file_str)}'!",
)
return None
choices = []
for i, c in enumerate(obj):
if isinstance(c, (str, int, float, bool)):
choices.append(str(c))
elif isinstance(c, list):
# we create an anonymous wildcard
choices.append(self.__create_anonymous_wildcard(full_path, key_parts, i, c))
elif isinstance(c, dict):
choices.append(self.__process_dict_choice(c, full_path, key_parts, i))
else:
log(
self.__logger,
self.__debug_level,
logging.WARNING,
f"Invalid choice {i+1} in wildcard '{escape_single_quotes('/'.join(key_parts))}' in file '{escape_single_quotes(file_str)}'!",
)
return choices
def __process_dict_choice(self, c: dict, full_path: Path | None, key_parts: list[str], i: int) -> dict:
"""
Process a dictionary choice.
Args:
c (dict): The dictionary choice.
full_path (Path | None): The path to the file, or None if from inline input.
key_parts (list[str]): The parts of the key.
i (int): The index of the choice.
Returns:
dict: The processed choice.
"""
if self.is_dict_wcdef_options(c):
return c
elif self.is_dict_choice_options(c):
# we assume it is a choice in object format
choice = c
choice_content = choice.get("content", choice.get("text", None))
if choice_content is not None and isinstance(choice_content, list):
# we create an anonymous wildcard
choice["content"] = self.__create_anonymous_wildcard(full_path, key_parts, i, choice_content)
if "text" in choice:
del choice["text"]
return choice
if len(c) == 1:
# we assume it is an anonymous wildcard with options
firstkey = list(c.keys())[0]
return self.__create_anonymous_wildcard(full_path, key_parts, i, c[firstkey], firstkey)
file_str = str(full_path) if full_path is not None else "input"
log(
self.__logger,
self.__debug_level,
logging.WARNING,
f"Invalid choice {i+1} in wildcard '{escape_single_quotes('/'.join(key_parts))}' in file '{escape_single_quotes(file_str)}'!",
)
return None
def __create_anonymous_wildcard(self, full_path: Path | None, key_parts, i, content, options=None):
"""
Create an anonymous wildcard.
Args:
full_path (Path | None): The path to the file that contains it, or None if from inline input.
key_parts (list[str]): The parts of the key.
i (int): The index of the wildcard.
content (object): The content of the wildcard.
options (str): The options for the choice where the wildcard is defined.
Returns:
str: The resulting value for the choice.
"""
new_parts = key_parts + [f"#ANON_{i}"]
self.__add_wildcard(content, full_path, new_parts)
value = f"__{'/'.join(new_parts)}__"
if options is not None:
value = f"{options}::{value}"
return value
def __add_wildcard(self, content: object, full_path: Path | None, external_key_parts: list[str]):
"""
Add a wildcard to the wildcards dictionary.
Args:
content (object): The content of the wildcard.
full_path (Path | None): The path to the file that contains it, or None if from inline input.
external_key_parts (list[str]): The parts of the key.
"""
file_str = str(full_path) if full_path is not None else "input"
def existing_file_str(wc):
return str(wc.file) if wc.file is not None else "input"
key_parts = external_key_parts.copy()
if isinstance(content, dict):
key_parts.pop()
keys = self.__get_wc_in_dict(content)
for key, obj in keys:
tmp_key_parts = key_parts.copy()
tmp_key_parts.extend(key.split("/"))
fullkey = "/".join(tmp_key_parts)
if self.wildcards.get(fullkey, None) is not None:
log(
self.__logger,
self.__debug_level,
logging.WARNING,
f"Duplicate wildcard '{escape_single_quotes(fullkey)}' in file '{escape_single_quotes(file_str)}' and '{escape_single_quotes(existing_file_str(self.wildcards[fullkey]))}'!",
)
else:
choices = self.__get_choices(obj, full_path, tmp_key_parts)
if choices is None:
log(
self.__logger,
self.__debug_level,
logging.WARNING,
f"Invalid wildcard '{escape_single_quotes(fullkey)}' in file '{escape_single_quotes(file_str)}'!",
)
elif fullkey.startswith("_"):
log(
self.__logger,
self.__debug_level,
logging.WARNING,
f"Invalid wildcard name '{escape_single_quotes(fullkey)}' in file '{escape_single_quotes(file_str)}'! (cannot start with underscore)",
)
else:
self.wildcards[fullkey] = PPPWildcard(full_path, fullkey, choices)
return
if isinstance(content, str):
content = [content]
elif isinstance(content, (int, float, bool)):
content = [str(content)]
if not isinstance(content, list):
log(
self.__logger,
self.__debug_level,
logging.WARNING,
f"Invalid wildcard in file '{escape_single_quotes(file_str)}'!",
)
return
fullkey = "/".join(key_parts)
if self.wildcards.get(fullkey, None) is not None:
log(
self.__logger,
self.__debug_level,
logging.WARNING,
f"Duplicate wildcard '{escape_single_quotes(fullkey)}' in file '{escape_single_quotes(file_str)}' and '{escape_single_quotes(existing_file_str(self.wildcards[fullkey]))}'!",
)
else:
choices = self.__get_choices(content, full_path, key_parts)
if choices is None:
log(
self.__logger,
self.__debug_level,
logging.WARNING,
f"Invalid wildcard '{escape_single_quotes(fullkey)}' in file '{escape_single_quotes(file_str)}'!",
)
elif fullkey.startswith("_"):
log(
self.__logger,
self.__debug_level,
logging.WARNING,
f"Invalid wildcard name '{escape_single_quotes(fullkey)}' in file '{escape_single_quotes(file_str)}'! (cannot start with underscore)",
)
else:
self.wildcards[fullkey] = PPPWildcard(full_path, fullkey, choices)
def __get_wildcards_in_structured_file(self, full_path: Path, base: Path):
"""
Get all wildcards in a structured file.
Args:
full_path (Path): The path to the file.
base (Path): The base path for the wildcards.
"""
external_key_parts = list(full_path.with_suffix("").relative_to(base).parts)
try:
with open(full_path, "r", encoding="utf-8") as file:
content = _YAML(typ='safe').load(file)
except: # pylint: disable=bare-except
log(
self.__logger,
self.__debug_level,
logging.WARNING,
f"Could not read file '{escape_single_quotes(str(full_path))}' with utf-8 encoding, trying windows-1252...",
)
with open(full_path, "r", encoding="windows-1252") as file:
content = _YAML(typ='safe').load(file)
self.__add_wildcard(content, full_path, external_key_parts)
def __get_wildcards_in_text_file(self, full_path: Path, base: Path):
"""
Get all wildcards in a text file.
Args:
full_path (Path): The path to the file.
base (Path): The base path for the wildcards.
"""
external_key_parts = list(full_path.with_suffix("").relative_to(base).parts)
try:
with open(full_path, "r", encoding="utf-8") as file:
text_content = map(lambda x: x.strip("\n\r"), file.readlines())
except: # pylint: disable=bare-except
log(
self.__logger,
self.__debug_level,
logging.WARNING,
f"Could not read file '{escape_single_quotes(str(full_path))}' with utf-8 encoding, trying windows-1252...",
)
with open(full_path, "r", encoding="windows-1252") as file:
text_content = map(lambda x: x.strip("\n\r"), file.readlines())
# First pass: drop blank lines and full-line comments.
text_content = list(filter(lambda x: x.strip() != "" and not x.strip().startswith("#"), text_content))
# Second pass: strip inline comments from lines that passed the first filter.
text_content = [x.split("#")[0].rstrip() if len(x.split("#")) > 1 else x for x in text_content]
self.__add_wildcard(text_content, full_path, external_key_parts)
def __get_wildcards_in_path(self, base: Path, path: Path):
"""
Get all wildcards in a path.
Args:
base (Path): The base path for the wildcards.
path (Path): The path (folder or file).
"""
if not path.exists():
log(
self.__logger,
self.__debug_level,
logging.WARNING,
f"Wildcard path '{escape_single_quotes(str(path))}' does not exist!",
)
return
if path.is_file():
self.__get_wildcards_in_file(base, path)
return
for child in path.iterdir():
if child.name.startswith("."):
continue
self.__get_wildcards_in_path(base, child)
def set_wildcard_default_filter(self, wildcard_key: str, filter_options: Optional[list[list[str]]]):
"""
Set the default filter for a wildcard.
Args:
wildcard_key (str): The key of the wildcard.
filter_options (list[list[str]]): The filter options.
"""
if filter_options is None:
if wildcard_key in self.__wildcard_default_filters:
del self.__wildcard_default_filters[wildcard_key]
else:
self.__wildcard_default_filters[wildcard_key] = filter_options
def get_wildcard_default_filter(self, wildcard_key: str) -> Optional[list[list[str]]]:
"""
Get the default filter for a wildcard.
Args:
wildcard_key (str): The key of the wildcard.
Returns:
Optional[list[list[str]]]: The filter options or None if not set.
"""
return self.__wildcard_default_filters.get(wildcard_key, None)
def reset_default_filters(self):
"""
Reset all default filters.
"""
self.__wildcard_default_filters = {}