diff --git a/.github/instructions/python-compat.instructions.md b/.github/instructions/python.instructions.md similarity index 84% rename from .github/instructions/python-compat.instructions.md rename to .github/instructions/python.instructions.md index 97ae083..205bf7d 100644 --- a/.github/instructions/python-compat.instructions.md +++ b/.github/instructions/python.instructions.md @@ -1,12 +1,14 @@ --- -description: "Use when writing, editing, or reviewing Python code. Enforces Python 3.10 compatibility — avoid syntax and stdlib features introduced in 3.11 or later." +description: "Use when writing, editing, or reviewing Python code." applyTo: "**/*.py" --- -# Python 3.10 Compatibility +# Python code + +## Python 3.10 Compatibility All Python code must be compatible with Python 3.10. Do not use language features or standard-library additions introduced in 3.11 or later. -## Forbidden (3.11+) +### Forbidden (3.11+) | Avoid | Use instead | |-------|-------------| @@ -18,7 +20,7 @@ All Python code must be compatible with Python 3.10. Do not use language feature | `except*` / `ExceptionGroup` | Not available; raise/catch normally | | `asyncio.TaskGroup`, `asyncio.timeout()` | `asyncio.gather()` / `asyncio.wait_for()` | -## Forbidden (3.12+) +### Forbidden (3.12+) | Avoid | Use instead | |-------|-------------| @@ -27,7 +29,7 @@ All Python code must be compatible with Python 3.10. Do not use language feature | `@typing.override` | Omit or use comment | | `itertools.batched()` | Manual chunking or `more-itertools` | -## Safe to use (available in 3.10) +### Safe to use (available in 3.10) - `match`/`case` structural pattern matching - `X | Y` union type syntax in annotations (e.g., `int | None`) @@ -35,3 +37,7 @@ All Python code must be compatible with Python 3.10. Do not use language feature - `list[int]`, `dict[str, int]` — built-in generic aliases - `zip(..., strict=True)` - `str.removeprefix()` / `str.removesuffix()` + +## Paths + +Use `pathlib.Path` for filesystem paths instead of `str` paths or `os.path`. diff --git a/__init__.py b/__init__.py index 8e7ec7c..2982309 100644 --- a/__init__.py +++ b/__init__.py @@ -6,9 +6,9 @@ """ import sys -import os +from pathlib import Path -sys.path.append(os.path.dirname(os.path.abspath(__file__))) +sys.path.append(str(Path(__file__).resolve().parent)) from .ppp_comfyui import ( PromptPostProcessorComfyUINode, diff --git a/install.py b/install.py index abdd67c..60292c7 100644 --- a/install.py +++ b/install.py @@ -1,6 +1,6 @@ -import os +from pathlib import Path -requirements_filename = os.path.join(os.path.dirname(os.path.realpath(__file__)), "requirements.txt") +requirements_filename = str(Path(__file__).resolve().parent / "requirements.txt") try: from modules.launch_utils import requirements_met, run_pip # A1111 diff --git a/ppp.py b/ppp.py index 9912de1..7ad377f 100644 --- a/ppp.py +++ b/ppp.py @@ -1,7 +1,7 @@ import dataclasses from enum import Enum import logging -import os +from pathlib import Path import re import time from typing import Any, Callable, Optional @@ -47,7 +47,7 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in """ version_str = "0.0.0" try: - pyproject_path = os.path.join(os.path.dirname(os.path.realpath(__file__)), "pyproject.toml") + pyproject_path = Path(__file__).resolve().parent / "pyproject.toml" with open(pyproject_path, "r", encoding="utf-8") as file: for line in file: if line.startswith("version = "): @@ -285,7 +285,7 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in def __load_config_and_detect(self, env_info: dict[str, Any]) -> HostConfig: """Loads config files, performs model detection, and returns the resolved host config.""" - default_config_file = os.path.join(os.path.dirname(os.path.realpath(__file__)), "ppp_config.yaml.defaults") + default_config_file = str(Path(__file__).resolve().parent / "ppp_config.yaml.defaults") try: with open(default_config_file, "r", encoding="utf-8") as f: default_raw: dict[str, Any] = yaml.safe_load(f) @@ -312,13 +312,13 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in import folder_paths # type: ignore user_dir = folder_paths.get_user_directory() - if user_dir and os.path.isdir(user_dir): - user_config_file = os.path.join(user_dir, "default", "ppp_config.yaml") + if user_dir and Path(user_dir).is_dir(): + user_config_file = str(Path(user_dir) / "default" / "ppp_config.yaml") except Exception: # pylint: disable=broad-exception-caught self.log(logging.WARNING, "Failed to get user directory for PPP config.") - if not user_config_file or not os.path.exists(user_config_file): - user_config_file = os.path.join(os.path.dirname(os.path.realpath(__file__)), "ppp_config.yaml") - if user_config_file and os.path.exists(user_config_file): + if not user_config_file or not Path(user_config_file).exists(): + user_config_file = str(Path(__file__).resolve().parent / "ppp_config.yaml") + if user_config_file and Path(user_config_file).exists(): with open(user_config_file, "r", encoding="utf-8") as f: user_raw = yaml.safe_load(f) user_cfg, _ = self.__parse_configuration(user_raw, "user configuration") @@ -374,7 +374,11 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in model_class = get_model_class_from_filename(model_name) if model_class: env_info["model_class"] = model_class - self.log(logging.DEBUG, f"Detected model class '{model_class}' from filename '{model_name}'", min_level=DEBUG_LEVEL.minimal) + self.log( + logging.DEBUG, + f"Detected model class '{model_class}' from filename '{model_name}'", + min_level=DEBUG_LEVEL.minimal, + ) app = env_info.get("app", "") for m in self.known_models: env_info["is_" + m] = False @@ -620,8 +624,8 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in vs.set_system("_sdfullname", model_filename) # deprecated vs.set_system("_modelfullname", model_filename) vs.set_system("_modelinfo", f"{self.state.env_info.get('model_class', '')}@{model_filename}") - vs.set_system("_sdname", os.path.basename(model_filename)) # deprecated - vs.set_system("_modelname", os.path.basename(model_filename)) + vs.set_system("_sdname", Path(model_filename).name) # deprecated + vs.set_system("_modelname", Path(model_filename).name) vs.set_system("_modelclass", self.state.env_info.get("model_class", "")) is_models = {} for model_name, model_type_and_substrings in self.variants_definitions.items(): diff --git a/ppp_classes.py b/ppp_classes.py index af7cf7f..b040a99 100644 --- a/ppp_classes.py +++ b/ppp_classes.py @@ -44,6 +44,7 @@ class ONWARNING_CHOICES(Enum): warn = "warn" stop = "stop" + # ------------------- Host configuration ------------------- AttentionOption = Literal["ok", "parentheses", "disable", "remove", "error"] @@ -172,8 +173,10 @@ class PPPConfig(BaseModel): raise ValueError("At least one of 'hosts' or 'models' must be specified") return self + # ------------------- State object ------------------- + @dataclass(frozen=True) class PPPStateOptions: """Options that can be set for prompt processing.""" @@ -221,6 +224,7 @@ class PPPStateOptions: object.__setattr__(self, "cup_merge_attention", False) object.__setattr__(self, "cup_remove_extranetwork_tags", False) + class CyclicalSamplerState: """Maintains the cycling position for '@' choice samplers across process_prompt calls.""" diff --git a/ppp_comfyui.py b/ppp_comfyui.py index 39636de..ce6750d 100644 --- a/ppp_comfyui.py +++ b/ppp_comfyui.py @@ -1,5 +1,6 @@ import logging import os +from pathlib import Path from typing import Any import folder_paths # type: ignore @@ -17,7 +18,7 @@ if __name__ == "__main__": raise SystemExit("This script must be run from ComfyUI") -def _resolve_wildcards_folders(override: str = "") -> list[str]: +def _resolve_wildcards_folders(override: str = "") -> list[Path]: """Return the resolved list of wildcard folder paths.""" folders_str = override if folders_str == "": @@ -33,13 +34,13 @@ def _resolve_wildcards_folders(override: str = "") -> list[str]: if folders_str == "": folders_str = os.getenv("WILDCARD_DIR", PPPWildcards.DEFAULT_WILDCARDS_FOLDER) return [ - (f if os.path.isabs(f) else os.path.abspath(os.path.join(folder_paths.models_dir, f))) + (Path(f) if Path(f).is_absolute() else (Path(folder_paths.models_dir) / f).resolve()) for f in folders_str.split(",") if f.strip() != "" ] -def _resolve_enmappings_folders(override: str = "") -> list[str]: +def _resolve_enmappings_folders(override: str = "") -> list[Path]: """Return the resolved list of extra-network mapping folder paths.""" folders_str = override if folders_str == "": @@ -51,7 +52,7 @@ def _resolve_enmappings_folders(override: str = "") -> list[str]: if folders_str == "": folders_str = os.getenv("EXTRANETWORKMAPPINGS_DIR", PPPExtraNetworkMappings.DEFAULT_ENMAPPINGS_FOLDER) return [ - (f if os.path.isabs(f) else os.path.abspath(os.path.join(folder_paths.models_dir, f))) + (Path(f) if Path(f).is_absolute() else (Path(folder_paths.models_dir) / f).resolve()) for f in folders_str.split(",") if f.strip() != "" ] @@ -67,7 +68,7 @@ class PromptPostProcessorComfyUINode: def __init__(self): lf = PromptPostProcessorLogFactory() self.logger = lf.log - grammar_filename = os.path.join(os.path.dirname(os.path.realpath(__file__)), "grammar.lark") + grammar_filename = Path(__file__).resolve().parent / "grammar.lark" with open(grammar_filename, "r", encoding="utf-8") as file: self.grammar_content = file.read() self.wildcards_obj = PPPWildcards(lf.log) @@ -438,7 +439,7 @@ class PromptPostProcessorComfyUINode: results = self.ppp.process_prompt(pos_prompt, neg_prompt, seed if seed is not None else 1) self.ppp.process_prompts_group_end() - # with open(os.path.join(os.path.dirname(os.path.realpath(__file__)), "logs", "last_prompts_comfyui.txt"), "w", encoding="utf-8") as f: + # with open(Path(__file__).parent / "logs" / "last_prompts_comfyui.txt", "w", encoding="utf-8") as f: # f.write(f"Seed: {seed if seed is not None else 1}\n") # f.write(f"In Positive: {pos_prompt}\n") # f.write(f"In Negative: {neg_prompt}\n") diff --git a/ppp_common.py b/ppp_common.py index 96e227f..fa30aaf 100644 --- a/ppp_common.py +++ b/ppp_common.py @@ -1,6 +1,6 @@ import ast import logging -import os +from pathlib import Path import re import textwrap import time @@ -92,7 +92,7 @@ def warn_or_stop(state: PPPState, is_negative: bool, message: str, e: Exception def load_grammar() -> str: # Process with lark (debug with https://www.lark-parser.org/ide/) - grammar_filename = os.path.join(os.path.dirname(os.path.realpath(__file__)), "grammar.lark") + grammar_filename = Path(__file__).resolve().parent / "grammar.lark" with open(grammar_filename, "r", encoding="utf-8") as file: grammar_content = file.read() return grammar_content @@ -197,6 +197,7 @@ def preprocess_grammar(grammar_content: str, options: dict[str, bool], logger: l ) return "\n".join(result_lines) + def get_model_class_from_filename(filename: str) -> str: try: import folder_paths # type: ignore @@ -208,9 +209,11 @@ def get_model_class_from_filename(filename: str) -> str: if not filename: return "" - full_path = folder_paths.get_full_path("checkpoints", filename) - if not full_path: - full_path = folder_paths.get_full_path("diffusion_models", filename) + full_path = ( + folder_paths.get_full_path("diffusion_models", filename) + or folder_paths.get_full_path("checkpoints", filename) + or folder_paths.get_full_path("unet", filename) + ) if not full_path or not full_path.lower().endswith((".safetensors", ".sft")): return "" try: diff --git a/ppp_enmappings.py b/ppp_enmappings.py index 4a0200f..366bf43 100644 --- a/ppp_enmappings.py +++ b/ppp_enmappings.py @@ -1,4 +1,3 @@ -import os from pathlib import Path from typing import Optional import logging @@ -35,12 +34,12 @@ class PPPENMapping: Attributes: kind (str): The kind of the extra network. name (str): The name of the extra network mapping. - file (str): The path to the file where the extranetwork mapping is defined. + file (Path | None): The path to the file where the extranetwork mapping is defined, or None if from inline input. variants (list[PPPENMappingVariant]): The processed variants of the extranetwork mapping. """ - def __init__(self, fullpath: str, kind: str, name: str, variants: list[dict]): - self.file: str = fullpath + def __init__(self, fullpath: Path | None, kind: str, name: str, variants: list[dict]): + self.file: Path | None = fullpath self.kind: str = kind self.name: str = name self.variants: list[PPPENMappingVariant] = [ @@ -55,7 +54,12 @@ class PPPENMapping: return hash(t) def __sizeof__(self): - return self.kind.__sizeof__() + self.name.__sizeof__() + self.file.__sizeof__() + self.variants.__sizeof__() + return ( + self.kind.__sizeof__() + + self.name.__sizeof__() + + (self.file.__sizeof__() if self.file is not None else 0) + + self.variants.__sizeof__() + ) class PPPExtraNetworkMappings: @@ -67,13 +71,13 @@ class PPPExtraNetworkMappings: """ DEFAULT_ENMAPPINGS_FOLDER = "extranetworkmappings" - LOCALINPUT_FILENAME = R"//INPUT\\" def __init__(self, logger=None): self.__logger: logging.Logger = logger self.__debug_level = DEBUG_LEVEL.none - self.__enmappings_folders = [] - self.__enmappings_files = {} + self.__enmappings_folders: list[Path] = [] + self.__enmappings_files: dict[Path, float] = {} + self.__local_enmappings_input_hash: int | None = None self.extranetwork_mappings: dict[str, PPPENMapping] = {} self.cached_mappings = {} @@ -89,25 +93,26 @@ class PPPExtraNetworkMappings: ) def refresh_extranetwork_mappings( - self, debug_level: DEBUG_LEVEL, enmappings_folders: Optional[list[str]], enmappings_input: str = None + self, + debug_level: DEBUG_LEVEL, + enmappings_folders: Optional[list[Path]], + enmappings_input: str = None, ): """ Initialize the extra network mappings. """ self.__debug_level = debug_level - self.__enmappings_folders = enmappings_folders or [] + self.__enmappings_folders = [Path(f) for f in (enmappings_folders or [])] # log(self.__logger, self.__debug_level, logging.INFO, "Refreshing extra network mappings...") # t1 = time.monotonic_ns() self.cached_mappings = {} for fullpath in list(self.__enmappings_files.keys()): - if fullpath != self.LOCALINPUT_FILENAME: - path = os.path.dirname(fullpath) - if not os.path.exists(fullpath) or not any( - Path(path).is_relative_to(folder) for folder in self.__enmappings_folders - ): - self.__remove_extranetwork_mappings_from_path(fullpath) - elif enmappings_input is None: + if not fullpath.exists() or not any( + fullpath.parent.is_relative_to(folder) for folder in self.__enmappings_folders + ): self.__remove_extranetwork_mappings_from_path(fullpath) + if enmappings_input is None and self.__local_enmappings_input_hash is not None: + self.__remove_extranetwork_mappings_from_input() if enmappings_folders is not None or enmappings_input is not None: if enmappings_folders is not None: for f in self.__enmappings_folders: @@ -117,60 +122,57 @@ class PPPExtraNetworkMappings: else: self.extranetwork_mappings = {} self.__enmappings_files = {} + self.__local_enmappings_input_hash = None # t2 = time.monotonic_ns() # log(self.__logger, self.__debug_level, logging.INFO, f"Extra network mappings refresh time: {(t2 - t1) / 1_000_000_000:.3f} seconds") - # def get_extranetwork_mappings(self, key: str) -> list[PPPENMapping]: - # """ - # Get all extra network mappings that match a key. - # - # Args: - # key (str): The key to match (kind:name). - # - # Returns: - # list: A list of all extra network mappings that match the key. - # """ - # keys = sorted(fnmatch.filter(self.extranetwork_mappings.keys(), key)) - # return [self.extranetwork_mappings[k] for k in keys] - - def __remove_extranetwork_mappings_from_path(self, full_path: str, debug=True): + def __remove_extranetwork_mappings_from_path(self, full_path: Path, debug=True): """ - Clear all extra network mappings in a file. + Clear all extra network mappings from a file. Args: - full_path (str): The path to the file. + full_path (Path): The path to the file. debug (bool): Whether to print debug messages or not. """ - last_modified_cached = self.__enmappings_files.get(full_path, None) # a time or a hash - if debug and last_modified_cached is not None: - if full_path == self.LOCALINPUT_FILENAME: - log(self.__logger, self.__debug_level, logging.DEBUG, "Removing extra network mappings from input") - else: - log( - self.__logger, - self.__debug_level, - logging.DEBUG, - f"Removing extra network mappings from file: {full_path}", - ) - if full_path in self.__enmappings_files.keys(): + if debug and full_path in self.__enmappings_files: + log( + self.__logger, + self.__debug_level, + logging.DEBUG, + f"Removing extra network mappings from file: {full_path}", + ) + if full_path in self.__enmappings_files: del self.__enmappings_files[full_path] for key in list(self.extranetwork_mappings.keys()): if self.extranetwork_mappings[key].file == full_path: del self.extranetwork_mappings[key] - def __get_extranetwork_mappings_in_file(self, full_path: str): + def __remove_extranetwork_mappings_from_input(self, debug=True): + """ + Clear all extra network mappings loaded from inline input. + + Args: + debug (bool): Whether to print debug messages or not. + """ + if debug and self.__local_enmappings_input_hash is not None: + log(self.__logger, self.__debug_level, logging.DEBUG, "Removing extra network mappings from input") + self.__local_enmappings_input_hash = None + for key in list(self.extranetwork_mappings.keys()): + if self.extranetwork_mappings[key].file is None: + del self.extranetwork_mappings[key] + + def __get_extranetwork_mappings_in_file(self, full_path: Path): """ Get all extra network mappings in a file. Args: - full_path (str): The path to the file. + full_path (Path): The path to the file. """ - last_modified = os.path.getmtime(full_path) + last_modified = full_path.stat().st_mtime last_modified_cached = self.__enmappings_files.get(full_path, None) if last_modified_cached is not None and last_modified == self.__enmappings_files[full_path]: return - filename = os.path.basename(full_path) - _, extension = os.path.splitext(filename) + extension = full_path.suffix if extension not in (".yaml", ".yml", ".json"): return self.__remove_extranetwork_mappings_from_path(full_path, False) @@ -192,11 +194,11 @@ class PPPExtraNetworkMappings: enmappings_input (str): The input string containing extra network mappings in yaml format. """ new_h = hash(enmappings_input) - h = self.__enmappings_files.get(self.LOCALINPUT_FILENAME, None) - if h == new_h: + if new_h == self.__local_enmappings_input_hash: return - self.__remove_extranetwork_mappings_from_path(self.LOCALINPUT_FILENAME, False) - if h is not None: + was_loaded = self.__local_enmappings_input_hash is not None + self.__remove_extranetwork_mappings_from_input(False) + if was_loaded: log(self.__logger, self.__debug_level, logging.DEBUG, "Updating extra network mappings from input") enmappings_input = enmappings_input.strip() if enmappings_input != "": @@ -211,23 +213,24 @@ class PPPExtraNetworkMappings: ) return if content is not None: - self.__add_extranetwork_mapping(content, self.LOCALINPUT_FILENAME) - self.__enmappings_files[self.LOCALINPUT_FILENAME] = new_h + self.__add_extranetwork_mapping(content, None) + self.__local_enmappings_input_hash = new_h - def __add_extranetwork_mapping(self, content: dict[str, dict[str, list[dict]]], full_path: str): + def __add_extranetwork_mapping(self, content: dict[str, dict[str, list[dict]]], full_path: Path | None): """ Add an extra network mapping to the extra network mappings dictionary. Args: content (object): The content of the extra network mapping. - full_path (str): The path to the file that contains it. + full_path (Path | None): The path to the file that contains it, or None if from inline input. """ + file_str = str(full_path) if full_path is not None else "input" if not isinstance(content, dict): log( self.__logger, self.__debug_level, logging.WARNING, - f"Invalid extra network mapping in file '{escape_single_quotes(full_path)}'!", + f"Invalid extra network mapping in file '{escape_single_quotes(file_str)}'!", ) return for kind, maps in content.items(): @@ -236,7 +239,7 @@ class PPPExtraNetworkMappings: self.__logger, self.__debug_level, logging.WARNING, - f"Invalid extra network mapping definition for '{escape_single_quotes(kind)}:*' in file '{escape_single_quotes(full_path)}'!", + f"Invalid extra network mapping definition for '{escape_single_quotes(kind)}:*' in file '{escape_single_quotes(file_str)}'!", ) else: for name, variants in maps.items(): @@ -246,32 +249,36 @@ class PPPExtraNetworkMappings: self.__logger, self.__debug_level, logging.WARNING, - f"Invalid extra network mapping definition for '{escape_single_quotes(key)}' in file '{escape_single_quotes(full_path)}'!", + f"Invalid extra network mapping definition for '{escape_single_quotes(key)}' in file '{escape_single_quotes(file_str)}'!", ) elif self.extranetwork_mappings.get(key, None) is not None: + f = ( + str(self.extranetwork_mappings[key].file) + if self.extranetwork_mappings[key].file is not None + else "input" + ) log( self.__logger, self.__debug_level, logging.WARNING, - f"Duplicate extra network mapping '{escape_single_quotes(key)}' in file '{escape_single_quotes(full_path)}' and '{escape_single_quotes(self.extranetwork_mappings[key].file)}'!", + f"Duplicate extra network mapping '{escape_single_quotes(key)}' in file '{escape_single_quotes(file_str)}' and '{escape_single_quotes(f)}'!", ) elif not isinstance(variants, list) or not all(isinstance(v, dict) for v in variants): log( self.__logger, self.__debug_level, logging.WARNING, - f"Invalid extra network mapping definition for '{escape_single_quotes(key)}' in file '{escape_single_quotes(full_path)}'!", + f"Invalid extra network mapping definition for '{escape_single_quotes(key)}' in file '{escape_single_quotes(file_str)}'!", ) else: self.extranetwork_mappings[key] = PPPENMapping(full_path, kind, name, variants) - def __get_extranetwork_mappings_in_structured_file(self, full_path): + def __get_extranetwork_mappings_in_structured_file(self, full_path: Path): """ Get all extra network mappings in a structured file. Args: - full_path (str): The path to the file. - base (str): The base path for the extra network mappings. + full_path (Path): The path to the file. """ try: try: @@ -282,7 +289,7 @@ class PPPExtraNetworkMappings: self.__logger, self.__debug_level, logging.WARNING, - f"Could not read file '{escape_single_quotes(full_path)}' with utf-8 encoding, trying windows-1252...", + 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.safe_load(file) @@ -292,29 +299,28 @@ class PPPExtraNetworkMappings: self.__logger, self.__debug_level, logging.ERROR, - f"Error reading extra network mappings from file '{escape_single_quotes(full_path)}': {e}", + f"Error reading extra network mappings from file '{escape_single_quotes(str(full_path))}': {e}", ) - def __get_extranetwork_mappings_in_path(self, path: str): + def __get_extranetwork_mappings_in_path(self, path: Path): """ Get all extra network mappings in a path. Args: - path (str): The path (folder or file). + path (Path): The path (folder or file). """ - if not os.path.exists(path): + if not path.exists(): log( self.__logger, self.__debug_level, logging.WARNING, - f"Extra network mappings path '{escape_single_quotes(path)}' does not exist!", + f"Extra network mappings path '{escape_single_quotes(str(path))}' does not exist!", ) return - if os.path.isfile(path): + if path.is_file(): self.__get_extranetwork_mappings_in_file(path) return - for filename in os.listdir(path): - full_path = os.path.abspath(os.path.join(path, filename)) - if os.path.basename(full_path).startswith("."): + for child in path.iterdir(): + if child.name.startswith("."): continue - self.__get_extranetwork_mappings_in_path(full_path) + self.__get_extranetwork_mappings_in_path(child) diff --git a/ppp_logging.py b/ppp_logging.py index 7566d1a..0914ae2 100644 --- a/ppp_logging.py +++ b/ppp_logging.py @@ -54,7 +54,7 @@ class PromptPostProcessorLogFactory: # pylint: disable=too-few-public-methods colored_record.levelname = f"{seq}{levelname:8s}{self.COLORS['RESET']}" return super().format(colored_record) - def __init__(self, filename = None): + def __init__(self, filename=None): """ Initializes the PromptPostProcessor class. @@ -98,7 +98,16 @@ class PromptPostProcessorLogCustomAdapter(logging.LoggerAdapter): """ return f"[PPP] {msg}", kwargs -def log(logger: logging.Logger, debug_level: DEBUG_LEVEL, kind: int, message: str, min_level: DEBUG_LEVEL | None = None, formatted: bool = True, exc_info: bool = False): + +def log( + logger: logging.Logger, + debug_level: DEBUG_LEVEL, + kind: int, + message: str, + min_level: DEBUG_LEVEL | None = None, + formatted: bool = True, + exc_info: bool = False, +): if logger: if min_level is None: if kind == logging.DEBUG: diff --git a/ppp_tree.py b/ppp_tree.py index d05bbbe..8cd0706 100644 --- a/ppp_tree.py +++ b/ppp_tree.py @@ -37,7 +37,12 @@ class TreeProcessor(lark.visitors.Interpreter): AccumulatedShell = namedtuple("AccumulatedShell", ["type", "data"]) NegTag = namedtuple("NegTag", ["start", "end", "content", "parameters", "shell"]) - def __init__(self, state: PPPState, rng: np.random.Generator, on_model_info_update: Optional[Callable[[], None]] = None): + def __init__( + self, + state: PPPState, + rng: np.random.Generator, + on_model_info_update: Optional[Callable[[], None]] = None, + ): super().__init__() self.state = state self.__on_model_info_update = on_model_info_update @@ -1235,16 +1240,14 @@ class TreeProcessor(lark.visitors.Interpreter): start_result = self.__result settable_sysvars = {"_modelfullname": "model_filename", "_modelclass": "model_class"} if self.state.variables.name_is_system(variable_name): - if (variable_name not in settable_sysvars and variable_name != "_modelinfo"): + if variable_name not in settable_sysvars and variable_name != "_modelinfo": self.warn_or_stop( f"Invalid variable name '{escape_single_quotes(variable_name)}' detected! System variables cannot be set." ) return app = self.state.env_info.get("app", "") if app not in (SUPPORTED_APPS.comfyui.value, SUPPORTED_APPS.tests.value): - self.warn_or_stop( - f"Setting '{escape_single_quotes(variable_name)}' is only supported in ComfyUI." - ) + self.warn_or_stop(f"Setting '{escape_single_quotes(variable_name)}' is only supported in ComfyUI.") return evaluated = self.__visit(content, restore_state=False, discard_content=True) if variable_name == "_modelinfo": diff --git a/ppp_wildcards.py b/ppp_wildcards.py index 605e40c..d66ca45 100644 --- a/ppp_wildcards.py +++ b/ppp_wildcards.py @@ -1,5 +1,4 @@ import fnmatch -import os from pathlib import Path from typing import Any, Optional import logging @@ -15,15 +14,15 @@ class PPPWildcard: Attributes: key (str): The key of the wildcard. - file (str): The path to the file where the wildcard is defined. + 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: str, key: str, choices: list[str]): + def __init__(self, fullpath: Path | None, key: str, choices: list[str]): self.key: str = key - self.file: str = fullpath + self.file: Path | None = fullpath self.unprocessed_choices: list[str] = choices self.choices: list[dict] = None self.options: dict = None @@ -35,7 +34,7 @@ class PPPWildcard: def __sizeof__(self): return ( self.key.__sizeof__() - + self.file.__sizeof__() + + (self.file.__sizeof__() if self.file is not None else 0) + self.unprocessed_choices.__sizeof__() + self.choices.__sizeof__() + self.options.__sizeof__() @@ -51,13 +50,13 @@ class PPPWildcards: """ DEFAULT_WILDCARDS_FOLDER = "wildcards" - LOCALINPUT_FILENAME = R"//INPUT\\" def __init__(self, logger=None): self.__logger: logging.Logger = logger self.__debug_level = DEBUG_LEVEL.none - self.__wildcards_folders = [] - self.__wildcard_files = {} + 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] = {} @@ -70,7 +69,7 @@ class PPPWildcards: def refresh_wildcards( self, debug_level: DEBUG_LEVEL, - wildcards_folders: Optional[list[str]], + wildcards_folders: Optional[list[Path]], wildcards_input: str = None, ): """ @@ -78,27 +77,26 @@ class PPPWildcards: """ self.reset_default_filters() self.__debug_level = debug_level - self.__wildcards_folders = wildcards_folders or [] + 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 fullpath != self.LOCALINPUT_FILENAME: - path = os.path.dirname(fullpath) - if not os.path.exists(fullpath) or not any( - Path(path).is_relative_to(folder) for folder in self.__wildcards_folders - ): - self.__remove_wildcards_from_path(fullpath) - elif wildcards_input is None: + 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 os.path.isdir(f) else os.path.dirname(f), f) + 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") @@ -134,46 +132,55 @@ class PPPWildcards: wc.append((prefix + str(key), obj)) return wc - def __remove_wildcards_from_path(self, full_path: str, debug=True): + def __remove_wildcards_from_path(self, full_path: Path, debug=True): """ - Clear all wildcards in a file. + Clear all wildcards from a file. Args: - full_path (str): The path to the file. + full_path (Path): The path to the file. debug (bool): Whether to print debug messages or not. """ - last_modified_cached = self.__wildcard_files.get(full_path, None) # a time or a hash - if debug and last_modified_cached is not None: - if full_path == self.LOCALINPUT_FILENAME: - log(self.__logger, self.__debug_level, logging.DEBUG, "Removing from memory wildcards from input") - else: - log( - self.__logger, - self.__debug_level, - logging.DEBUG, - f"Removing from memory wildcards from file: {full_path}", - ) - if full_path in self.__wildcard_files.keys(): + 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 __get_wildcards_in_file(self, base, full_path: str): + 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 (str): The base path for the wildcards. - full_path (str): The path to the file. + base (Path): The base path for the wildcards. + full_path (Path): The path to the file. """ try: - last_modified = os.path.getmtime(full_path) + 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 - filename = os.path.basename(full_path) - _, extension = os.path.splitext(filename) + extension = full_path.suffix if extension not in (".txt", ".json", ".yaml", ".yml"): return self.__remove_wildcards_from_path(full_path, False) @@ -189,7 +196,7 @@ class PPPWildcards: self.__logger, self.__debug_level, logging.ERROR, - f"Error reading wildcard file '{escape_single_quotes(full_path)}': {e}", + f"Error reading wildcard file '{escape_single_quotes(str(full_path))}': {e}", ) def __get_wildcards_in_input(self, wildcards_input: str): @@ -201,11 +208,11 @@ class PPPWildcards: """ try: new_h = hash(wildcards_input) - h = self.__wildcard_files.get(self.LOCALINPUT_FILENAME, None) - if h == new_h: + if new_h == self.__local_input_hash: return - self.__remove_wildcards_from_path(self.LOCALINPUT_FILENAME, False) - if h is not None: + 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 != "": @@ -215,8 +222,8 @@ class PPPWildcards: 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, self.LOCALINPUT_FILENAME, [self.LOCALINPUT_FILENAME]) - self.__wildcard_files[self.LOCALINPUT_FILENAME] = new_h + 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}") @@ -262,13 +269,13 @@ class PPPWildcards: """ return all(k in ["command", "labels", "weight", "if", "content", "text"] for k in d.keys()) - def __get_choices(self, obj: object, full_path: str, key_parts: list[str]) -> list: + 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 (str): path to the file where the wildcard is defined + 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: @@ -280,12 +287,13 @@ class PPPWildcards: 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(full_path)}'!", + f"Invalid format in wildcard '{escape_single_quotes('/'.join(key_parts))}' in file '{escape_single_quotes(file_str)}'!", ) return None choices = [] @@ -302,17 +310,17 @@ class PPPWildcards: 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(full_path)}'!", + 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: str, key_parts: list[str], i: int) -> dict: + 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 (str): The path to the file. + 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. @@ -335,20 +343,21 @@ class PPPWildcards: # 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(full_path)}'!", + 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, key_parts, i, content, options=None): + def __create_anonymous_wildcard(self, full_path: Path | None, key_parts, i, content, options=None): """ Create an anonymous wildcard. Args: - full_path (str): The path to the file that contains it. + 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. @@ -364,15 +373,20 @@ class PPPWildcards: value = f"{options}::{value}" return value - def __add_wildcard(self, content: object, full_path: str, external_key_parts: list[str]): + 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 (str): The path to the file that contains it. + 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() @@ -386,7 +400,7 @@ class PPPWildcards: self.__logger, self.__debug_level, logging.WARNING, - f"Duplicate wildcard '{escape_single_quotes(fullkey)}' in file '{escape_single_quotes(full_path)}' and '{escape_single_quotes(self.wildcards[fullkey].file)}'!", + 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) @@ -395,14 +409,14 @@ class PPPWildcards: self.__logger, self.__debug_level, logging.WARNING, - f"Invalid wildcard '{escape_single_quotes(fullkey)}' in file '{escape_single_quotes(full_path)}'!", + 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(full_path)}'! (cannot start with underscore)", + 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) @@ -416,7 +430,7 @@ class PPPWildcards: self.__logger, self.__debug_level, logging.WARNING, - f"Invalid wildcard in file '{escape_single_quotes(full_path)}'!", + f"Invalid wildcard in file '{escape_single_quotes(file_str)}'!", ) return fullkey = "/".join(key_parts) @@ -425,7 +439,7 @@ class PPPWildcards: self.__logger, self.__debug_level, logging.WARNING, - f"Duplicate wildcard '{escape_single_quotes(fullkey)}' in file '{escape_single_quotes(full_path)}' and '{escape_single_quotes(self.wildcards[fullkey].file)}'!", + 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) @@ -434,28 +448,27 @@ class PPPWildcards: self.__logger, self.__debug_level, logging.WARNING, - f"Invalid wildcard '{escape_single_quotes(fullkey)}' in file '{escape_single_quotes(full_path)}'!", + 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(full_path)}'! (cannot start with underscore)", + 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, base): + def __get_wildcards_in_structured_file(self, full_path: Path, base: Path): """ Get all wildcards in a structured file. Args: - full_path (str): The path to the file. - base (str): The base path for the wildcards. + full_path (Path): The path to the file. + base (Path): The base path for the wildcards. """ - external_key: str = os.path.relpath(os.path.splitext(full_path)[0], base) - external_key_parts = external_key.split(os.sep) + 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.safe_load(file) @@ -464,22 +477,21 @@ class PPPWildcards: self.__logger, self.__debug_level, logging.WARNING, - f"Could not read file '{escape_single_quotes(full_path)}' with utf-8 encoding, trying windows-1252...", + 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.safe_load(file) self.__add_wildcard(content, full_path, external_key_parts) - def __get_wildcards_in_text_file(self, full_path, base): + def __get_wildcards_in_text_file(self, full_path: Path, base: Path): """ Get all wildcards in a text file. Args: - full_path (str): The path to the file. - base (str): The base path for the wildcards. + full_path (Path): The path to the file. + base (Path): The base path for the wildcards. """ - external_key: str = os.path.relpath(os.path.splitext(full_path)[0], base) - external_key_parts = external_key.split(os.sep) + 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()) @@ -488,7 +500,7 @@ class PPPWildcards: self.__logger, self.__debug_level, logging.WARNING, - f"Could not read file '{escape_single_quotes(full_path)}' with utf-8 encoding, trying windows-1252...", + 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()) @@ -496,30 +508,29 @@ class PPPWildcards: 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: str, path: str): + def __get_wildcards_in_path(self, base: Path, path: Path): """ Get all wildcards in a path. Args: - base (str): The base path for the wildcards. - path (str): The path (folder or file). + base (Path): The base path for the wildcards. + path (Path): The path (folder or file). """ - if not os.path.exists(path): + if not path.exists(): log( self.__logger, self.__debug_level, logging.WARNING, - f"Wildcard path '{escape_single_quotes(path)}' does not exist!", + f"Wildcard path '{escape_single_quotes(str(path))}' does not exist!", ) return - if os.path.isfile(path): + if path.is_file(): self.__get_wildcards_in_file(base, path) return - for filename in os.listdir(path): - full_path = os.path.abspath(os.path.join(path, filename)) - if os.path.basename(full_path).startswith("."): + for child in path.iterdir(): + if child.name.startswith("."): continue - self.__get_wildcards_in_path(base, full_path) + self.__get_wildcards_in_path(base, child) def set_wildcard_default_filter(self, wildcard_key: str, filter_options: Optional[list[list[str]]]): """ @@ -550,4 +561,4 @@ class PPPWildcards: """ Reset all default filters. """ - self.__wildcard_default_filters = {} \ No newline at end of file + self.__wildcard_default_filters = {} diff --git a/scripts/ppp_script.py b/scripts/ppp_script.py index 7c15ee7..44f2670 100644 --- a/scripts/ppp_script.py +++ b/scripts/ppp_script.py @@ -61,9 +61,10 @@ class PromptPostProcessorA1111Script(scripts.Script): Returns: None """ + super().__init__() self.instance_index = self.increment_instance_count() self.name = PromptPostProcessor.NAME - grammar_filename = os.path.join(os.path.dirname(os.path.realpath(__file__)), "../grammar.lark") + grammar_filename = Path(__file__).resolve().parent.parent / "grammar.lark" with open(grammar_filename, "r", encoding="utf-8") as file: self.grammar_content = file.read() self.ppp_logger = None @@ -311,7 +312,7 @@ class PromptPostProcessorA1111Script(scripts.Script): if wc_wildcards_folders == "": wc_wildcards_folders = os.getenv("WILDCARD_DIR", PPPWildcards.DEFAULT_WILDCARDS_FOLDER) wildcards_folders = [ - (f if os.path.isabs(f) else os.path.abspath(os.path.join(models_path, f))) + (Path(f) if Path(f).is_absolute() else (Path(models_path) / f).resolve()) for f in wc_wildcards_folders.split(",") if f.strip() != "" ] @@ -322,7 +323,7 @@ class PromptPostProcessorA1111Script(scripts.Script): PPPExtraNetworkMappings.DEFAULT_ENMAPPINGS_FOLDER, ) enmappings_folders = [ - (f if os.path.isabs(f) else os.path.abspath(os.path.join(models_path, f))) + (Path(f) if Path(f).is_absolute() else (Path(models_path) / f).resolve()) for f in en_mappings_folders.split(",") if f.strip() != "" ] @@ -383,11 +384,11 @@ class PromptPostProcessorA1111Script(scripts.Script): regular_type = "regular" rpr: list[str] = getattr(p, "all_prompts", None) rnr: list[str] = getattr(p, "all_negative_prompts", None) - regular_exists = rpr is not None and rnr is not None + regular_exists = bool(rpr) and bool(rnr) hiresfix_type = "hiresfix" rph: list[str] = getattr(p, "all_hr_prompts", None) rnh: list[str] = getattr(p, "all_hr_negative_prompts", None) - hiresfix_exists = rph is not None and rnh is not None + hiresfix_exists = bool(rph) and bool(rnh) for i in range(len(calculated_seeds)): if regular_exists: prompts_list[(regular_type, i)] = None @@ -461,7 +462,7 @@ class PromptPostProcessorA1111Script(scripts.Script): prompts_list[(prompttype, typeindex)] = cached ppp.process_prompts_group_end() - # with open(os.path.join(os.path.dirname(os.path.realpath(__file__)), "..", "logs", f"last_prompts_{app.value}.txt"), "w", encoding="utf-8") as f: + # with open(Path(__file__).parent.parent / "logs" / f"last_prompts_{app.value}.txt", "w", encoding="utf-8") as f: # for (prompttype, typeindex), (posp, negp) in prompts_list.items(): # f.write(f"Key: {prompttype}[{typeindex}]\n") # f.write(f"Seed: {calculated_seeds[typeindex]}\n") diff --git a/tests/base_tests.py b/tests/base_tests.py index ca78914..f469bc4 100644 --- a/tests/base_tests.py +++ b/tests/base_tests.py @@ -1,7 +1,7 @@ from dataclasses import replace import difflib -import os import logging +from pathlib import Path from typing import Any, NamedTuple, Optional import unittest import datetime @@ -91,8 +91,8 @@ class TestPromptPostProcessorBase(unittest.TestCase): self.wildcards_obj.refresh_wildcards( DEBUG_LEVEL.full, [ - os.path.abspath(os.path.join(os.path.dirname(__file__), "wildcards")), - os.path.abspath(os.path.join(os.path.dirname(__file__), "wildcards2")), + Path(__file__).parent / "wildcards", + Path(__file__).parent / "wildcards2", ], """ yaml_input: @@ -104,11 +104,11 @@ class TestPromptPostProcessorBase(unittest.TestCase): ) self.extranetwork_maps_obj.refresh_extranetwork_mappings( DEBUG_LEVEL.full, - [os.path.abspath(os.path.join(os.path.dirname(__file__), "enmappings"))], + [Path(__file__).parent / "enmappings"], """ """, ) - grammar_filename = os.path.join(os.path.dirname(os.path.realpath(__file__)), "../grammar.lark") + grammar_filename = Path(__file__).resolve().parent.parent / "grammar.lark" with open(grammar_filename, "r", encoding="utf-8") as file: self.grammar_content = file.read() @@ -201,8 +201,8 @@ class TestPromptPostProcessorBase(unittest.TestCase): interrupted: bool = False, combinatorial: bool = False, combinatorial_limit: int = 0, - specific_wc_folders: Optional[list[str]] = None, - specific_em_folders: Optional[list[str]] = None, + specific_wc_folders: Optional[list[Path]] = None, + specific_em_folders: Optional[list[Path]] = None, ): """ Process the prompt and compare the results with the expected prompts. @@ -215,8 +215,8 @@ class TestPromptPostProcessorBase(unittest.TestCase): interrupted (bool, optional): The interrupted flag. Defaults to False. combinatorial (bool, optional): The combinatorial flag. Defaults to False. combinatorial_limit (int, optional): The combinatorial limit. Defaults to 0. - specific_wc_folders (Optional[list[str]], optional): A list of specific wildcard folders to refresh. Defaults to None. - specific_em_folders (Optional[list[str]], optional): A list of specific extranetwork mapping folders to refresh. Defaults to None. + specific_wc_folders (Optional[list[Path]], optional): A list of specific wildcard folders to refresh. Defaults to None. + specific_em_folders (Optional[list[Path]], optional): A list of specific extranetwork mapping folders to refresh. Defaults to None. Returns: None