Files
acorderob-sd-webui-prompt-p…/ppp_common.py
T
Antonio Cordero Balcazar a8a34f47e9 * Independent support for Forge Neo (from Classic).
* Styles (A1111/SD.Next) to Wildcard conversion tool.
* Additional models support and NoobAI variant for SDXL.
* Detection of new models supported but missing from user config file.
2026-05-19 18:05:12 +02:00

307 lines
12 KiB
Python

import ast
import csv
from functools import reduce
import logging
from pathlib import Path
import re
import textwrap
import time
import lark
import yaml
from ppp_logging import log
from ppp_classes import ONWARNING_CHOICES, PPPInterrupt, PPPState
from ppp_utils import escape_single_quotes, format_output
def parse_prompt(
state: PPPState,
prompt_description: str,
prompt: str,
parser: lark.Lark,
raise_parsing_error: bool = False,
):
"""
Parses a prompt using the specified parser.
Args:
prompt_description (str): The description of the prompt.
prompt (str): The prompt to be parsed.
parser (lark.Lark): The parser to be used.
raise_parsing_error (bool): Whether to raise a parsing error.
Returns:
Tree: The parsed prompt.
"""
t1 = time.monotonic_ns()
parsed_prompt = None
try:
log(
state.logger,
state.options.debug_level,
logging.DEBUG,
f"Parsing {prompt_description}: '{escape_single_quotes(prompt)}'",
)
parsed_prompt = parser.parse(prompt)
# we store the contents so we can use them later even if the meta position is not valid anymore
if isinstance(parsed_prompt, lark.Tree):
for n in parsed_prompt.iter_subtrees():
if isinstance(n, lark.Tree):
if n.meta.empty:
n.meta.content = ""
else:
n.meta.content = prompt[n.meta.start_pos : n.meta.end_pos]
except lark.exceptions.UnexpectedInput:
if raise_parsing_error:
raise
log(
state.logger,
state.options.debug_level,
logging.ERROR,
f"Parsing failed on prompt!: {escape_single_quotes(prompt)}",
)
t2 = time.monotonic_ns()
log(
state.logger,
state.options.debug_level,
logging.DEBUG,
f"Parse {prompt_description} time: {(t2 - t1) / 1_000_000_000:.3f} seconds",
)
if parsed_prompt:
log(
state.logger,
state.options.debug_level,
logging.DEBUG,
"Tree:\n"
+ textwrap.indent(
re.sub(r"\n$", "", (parsed_prompt.pretty() if isinstance(parsed_prompt, lark.Tree) else parsed_prompt)),
" ",
),
formatted=False,
)
return parsed_prompt
def warn_or_stop(state: PPPState, is_negative: bool, message: str, e: Exception = None):
INVALID_CONTENT_STOP = "INVALID CONTENT! {0}\nBREAK "
if state.options.on_warning == ONWARNING_CHOICES.stop:
raise PPPInterrupt(
message,
INVALID_CONTENT_STOP.format(message) if not is_negative else "",
INVALID_CONTENT_STOP.format(message) if is_negative else "",
) from e
log(state.logger, state.options.debug_level, logging.WARNING, format_output(message))
def load_grammar() -> str:
# Process with lark (debug with https://www.lark-parser.org/ide/)
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
def preprocess_grammar(grammar_content: str, options: dict[str, bool], logger: logging.Logger, debug_level: int) -> str:
"""
Preprocesses the grammar content to handle conditional compilation directives.
Args:
grammar_content (str): The raw grammar content.
options (dict[str,bool]): Options for preprocessing.
logger (logging.Logger): The logger object.
debug_level (int): The debug level for logging.
Returns:
str: The preprocessed grammar content.
"""
lines = grammar_content.split("\n")
result_lines = []
skip_current_block = []
all_blocks_skipped = []
def eval_bool_expr(expr: str, constants: dict[str, bool]) -> bool:
"""
Evaluates a boolean expression using known constants.
Supports: and, or, not, parentheses, and named constants.
Args:
expr (str): The boolean expression to evaluate.
constants (dict[str, bool]): A dictionary of constant values.
Returns:
bool: The result of the evaluated expression.
"""
tree = ast.parse(expr, mode="eval")
def _eval(node) -> bool:
if isinstance(node, ast.Expression):
return _eval(node.body)
if isinstance(node, ast.BoolOp):
if isinstance(node.op, ast.And):
return all(_eval(v) for v in node.values)
if isinstance(node.op, ast.Or):
return any(_eval(v) for v in node.values)
if isinstance(node, ast.UnaryOp) and isinstance(node.op, ast.Not):
return not _eval(node.operand)
if isinstance(node, ast.Name):
return bool(constants[node.id]) # raises KeyError for unknown names
if isinstance(node, ast.Constant) and isinstance(node.value, bool):
return node.value
raise ValueError(f"Unsupported construct: {ast.dump(node)}")
return _eval(tree)
for line in lines:
stripped_line = line.strip()
if stripped_line.startswith("//#if"):
# Extract condition from the #if directive
conditions = stripped_line[5:].strip()
# Evaluate the conditions
skip_current_block.append(not eval_bool_expr(conditions, options))
all_blocks_skipped.append(skip_current_block[-1])
elif stripped_line.startswith("//#elif"):
if not skip_current_block:
log(logger, debug_level, logging.WARNING, "Unmatched //#elif directive found in grammar content.")
elif all_blocks_skipped[-1]:
# Extract condition from the #elif directive
conditions = stripped_line[7:].strip()
# Evaluate the conditions
skip_current_block[-1] = not eval_bool_expr(conditions, options)
if not skip_current_block[-1]:
all_blocks_skipped[-1] = False
else:
skip_current_block[-1] = True
elif stripped_line.startswith("//#else"):
if not skip_current_block:
log(logger, debug_level, logging.WARNING, "Unmatched //#else directive found in grammar content.")
elif all_blocks_skipped[-1]:
skip_current_block[-1] = False
else:
skip_current_block[-1] = True
elif stripped_line.startswith("//#endif"):
if not skip_current_block:
log(logger, debug_level, logging.WARNING, "Unmatched //#endif directive found in grammar content.")
else:
skip_current_block.pop()
all_blocks_skipped.pop()
elif stripped_line.startswith("//#"):
log(
logger,
debug_level,
logging.WARNING,
f"Unrecognized directive found in grammar content: {stripped_line}",
)
elif not any(skip_current_block):
# Include the line if we're not skipping any current block
result_lines.append(stripped_line)
# Check for unclosed blocks at the end
if skip_current_block:
raise PPPInterrupt(
f"Found {len(skip_current_block)} unclosed conditional directive(s) at the end of the grammar file"
)
return "\n".join(result_lines)
def get_model_class_from_filename(filename: str) -> str:
try:
import folder_paths # type: ignore
import comfy.utils # type: ignore
import comfy.model_detection as model_detection # type: ignore
except ImportError:
return ""
import json
if not filename:
return ""
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:
header_bytes = comfy.utils.safetensors_header(full_path)
if header_bytes is None:
return ""
header = json.loads(header_bytes)
# Build a mock state dict - detection only needs key names and shapes
class _ShapeProxy:
def __init__(self, shape):
self.shape = shape
def __getitem__(self, i):
return self.shape[i] # shape[i] access
mock_sd = {k: _ShapeProxy(v["shape"]) for k, v in header.items() if k != "__metadata__" and "shape" in v}
prefix = model_detection.unet_prefix_from_state_dict(mock_sd)
config = model_detection.model_config_from_unet(mock_sd, prefix)
return config.__class__.__name__ if config else ""
except Exception: # pylint: disable=broad-except
return ""
def sanitize_wc_name(name: str) -> str:
# Remove invalid characters
return re.sub(r"[^a-zA-Z0-9-_]+", "", re.sub(r"_{2,}", "_", name.replace(" ", "_")))
def convert_a1111_styles_to_wildcard(inp: Path, out: Path):
"""
Converts styles from A1111 format to wildcard format in a YAML file.
Args:
inp (Path): The input CSV file path.
out (Path): The output YAML file path.
"""
wildcards = {}
with open(inp, "r", encoding="utf-8-sig") as f:
for reg in csv.reader(f):
name = reg[0].strip()
if name.lower() == "name":
continue
positive = reg[1].strip()
negative = reg[2].strip()
if negative:
positive = f"{positive}<ppp:stn>{negative}<ppp:/stn>"
name = sanitize_wc_name(name)
wildcards[name] = positive
if not wildcards:
raise RuntimeError(f"No styles found in {inp} to convert to wildcards.")
with open(out, "w", encoding="utf-8-sig") as f:
f.write(f"# Original names may contain characters that are replaced in the output.\n# Converted from {inp}\n")
yaml.dump(wildcards, f, allow_unicode=True)
def convert_sdnext_styles_to_wildcard(inp: Path, out: Path):
"""
Converts styles from SD.Next format to wildcard format in a YAML file.
Args:
inp (Path): The input folder path containing style json files or a single json file.
out (Path): The output YAML file path.
"""
wildcards = {}
files = inp.glob("*.json") if inp.is_dir() else [inp]
for file in files:
with open(file, "r", encoding="utf-8-sig") as f:
data = yaml.safe_load(f)
if not isinstance(data, list):
continue
wildcards[file] = {}
for style in data:
name = style.get("name", "").strip()
positive = style.get("prompt", "").strip()
negative = style.get("negative", "").strip()
# extra = style.get("extra", "").strip()
if negative:
positive = f"{positive}<ppp:stn>{negative}<ppp:/stn>"
name = sanitize_wc_name(name)
wildcards[file][name] = positive
if not reduce(lambda acc, d: acc or bool(d), wildcards.values(), False):
raise RuntimeError(f"No styles found in {inp} to convert to wildcards.")
with open(out, "w", encoding="utf-8-sig") as f:
f.write("# Original names may contain characters that are replaced in the output.\n")
for name, wcs in wildcards.items():
f.write(f"# Converted from {name}\n")
yaml.dump(wcs, f, allow_unicode=True)