Latest changes
This commit is contained in:
+122
-13
@@ -1,3 +1,6 @@
|
||||
import itertools
|
||||
|
||||
import comfy.utils
|
||||
import nodes
|
||||
|
||||
from . import evaluator
|
||||
@@ -42,10 +45,13 @@ class SrlFormatString:
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"format": ("STRING", {
|
||||
"multiline": False,
|
||||
"default": "first input via str(): {}, second input via repr(): {!r}, third input by index: {2}, fifth input by name: {in4}",
|
||||
}),
|
||||
"format": (
|
||||
"STRING",
|
||||
{
|
||||
"multiline": False,
|
||||
"default": "first input via str(): {}, second input via repr(): {!r}, third input by index: {2}",
|
||||
},
|
||||
),
|
||||
},
|
||||
"optional": {
|
||||
"in0": (any_typ,),
|
||||
@@ -61,23 +67,23 @@ class SrlFormatString:
|
||||
CATEGORY = "utils"
|
||||
|
||||
def doit(self, format, **kwargs):
|
||||
# Allow referencing arguments both by name and index.
|
||||
return (evaluator.safe_vformat(format, list(kwargs.values()), kwargs),)
|
||||
return (evaluator.safe_vformat(format, list(kwargs.values()), {}),)
|
||||
|
||||
|
||||
class SrlNumExpr:
|
||||
"""Evaluate a numerical expression safely."""
|
||||
|
||||
def __init__(self):
|
||||
self.last_expr = None
|
||||
self.last_code = None
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"expr": ("STRING", {"multiline": True}),
|
||||
"expr": ("STRING", {"multiline": True, "dynamicPrompts": False}),
|
||||
},
|
||||
"optional": {
|
||||
"b0": ("BOOLEAN", {"forceInput": True}),
|
||||
"b1": ("BOOLEAN", {"forceInput": True}),
|
||||
"b2": ("BOOLEAN", {"forceInput": True}),
|
||||
"i0": ("INT", {"forceInput": True}),
|
||||
"i1": ("INT", {"forceInput": True}),
|
||||
"i2": ("INT", {"forceInput": True}),
|
||||
@@ -87,10 +93,18 @@ class SrlNumExpr:
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("BOOL", "INT", "FLOAT")
|
||||
RETURN_TYPES = ("BOOLEAN", "INT", "FLOAT")
|
||||
FUNCTION = "doit"
|
||||
CATEGORY = "utils"
|
||||
|
||||
def doit(self, expr, **kwargs):
|
||||
pass
|
||||
if expr == self.last_expr:
|
||||
code = self.last_code
|
||||
else:
|
||||
code = self.last_code = evaluator.safe_compile(expr)
|
||||
|
||||
res = evaluator.safe_eval(code, kwargs)
|
||||
return (bool(res), int(res), float(res))
|
||||
|
||||
|
||||
class SrlFilterImageList:
|
||||
@@ -110,6 +124,7 @@ class SrlFilterImageList:
|
||||
INPUT_IS_LIST = True
|
||||
OUTPUT_IS_LIST = (True, True)
|
||||
FUNCTION = "doit"
|
||||
CATEGORY = "utils"
|
||||
|
||||
def doit(self, images, keep):
|
||||
im1, im2 = itertools.tee(zip(images, keep))
|
||||
@@ -133,6 +148,7 @@ class SrlCountSegs:
|
||||
|
||||
RETURN_TYPES = ("INT",)
|
||||
FUNCTION = "doit"
|
||||
CATEGORY = "utils"
|
||||
|
||||
def doit(self, segs, ignore_none):
|
||||
if segs is None:
|
||||
@@ -141,7 +157,96 @@ class SrlCountSegs:
|
||||
else:
|
||||
return 0
|
||||
else:
|
||||
return len(segs[1])
|
||||
return (len(segs[1]),)
|
||||
|
||||
|
||||
class SrlScaleImage:
|
||||
"""Scale an image under specified conditions."""
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"image": ("IMAGE",),
|
||||
"what": ("width", "height", "both", "factor", "longest", "shortest", "mebipixels"),
|
||||
"how": ("up", "down", "both"),
|
||||
"width": ("INT", {"default": 0, "min": 0, "max": nodes.MAX_RESOLUTION, "step": 1}),
|
||||
"height": ("INT", {"default": 0, "min": 0, "max": nodes.MAX_RESOLUTION, "step": 1}),
|
||||
"factor": ("FLOAT",),
|
||||
"longest_shortest": ("INT", {"default": 0, "min": 0, "max": nodes.MAX_RESOLUTION, "step": 1}),
|
||||
"mebipixels": ("FLOAT",),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("IMAGE", "INT", "INT", "FLOAT", "FLOAT", "BOOLEAN", "BOOLEAN")
|
||||
RETURN_NAMES = ("image", "width", "height", "width_factor", "height_factor", "did_scale", "did_scale_up")
|
||||
FUNCTION = "doit"
|
||||
CATEGORY = "images"
|
||||
|
||||
def scale_dim(self, how, idim, odim, width, height):
|
||||
scaled = scaled_up = False
|
||||
factor = 1.0
|
||||
|
||||
if how in ("up", "both") and idim < odim:
|
||||
scaled = scaled_up = True
|
||||
factor = odim / idim
|
||||
elif how in ("down", "both") and idim > odim:
|
||||
scaled = True
|
||||
factor = odim / idim
|
||||
|
||||
if scaled:
|
||||
width *= factor
|
||||
height *= factor
|
||||
|
||||
return width, height, factor, scaled, scaled_up
|
||||
|
||||
def doit(self, image, what, how, width, height, factor, longest_shortest, mebipixels, ratio):
|
||||
samples = image.movedim(-1, 1)
|
||||
w, h = samples.shape[2:4]
|
||||
scaled = scaled_up = False
|
||||
new_width = w
|
||||
new_height = h
|
||||
width_factor = height_factor = 1.0
|
||||
|
||||
match what:
|
||||
case "factor":
|
||||
new_width = w * factor
|
||||
new_height = h * factor
|
||||
scaled = factor != 1.0
|
||||
scaled_up = factor > 1.0
|
||||
width_factor = height_factor = factor
|
||||
case "both":
|
||||
new_width = width
|
||||
new_height = height
|
||||
width_factor = new_width / w
|
||||
height_factor = new_height / h
|
||||
scaled = new_width != width or new_height != height
|
||||
scaled_up = new_width > width
|
||||
case _:
|
||||
match what:
|
||||
case "width":
|
||||
idim = w
|
||||
odim = width
|
||||
case "height":
|
||||
idim = h
|
||||
odim = height
|
||||
case "longest":
|
||||
idim = max(width, height)
|
||||
odim = longest_shortest
|
||||
case "shortest":
|
||||
idim = min(width, height)
|
||||
odim = longest_shortest
|
||||
case "mebipixels":
|
||||
idim = width * height / 1048576
|
||||
odim = mebipixels * 1048576
|
||||
|
||||
new_width, new_height, ofactor, scaled, scaled_up = self.scale_dim(how, idim, odim, width, height)
|
||||
width_factor = height_factor = ofactor
|
||||
|
||||
if scaled:
|
||||
samples = comfy.utils.common_upscale(samples, width, height, "lanczos", "disabled")
|
||||
|
||||
return (samples.movedim(1, -1), new_width, new_height, width_factor, height_factor, scaled, scaled_up)
|
||||
|
||||
|
||||
# A dictionary that contains all nodes you want to export with their names
|
||||
@@ -150,7 +255,9 @@ NODE_CLASS_MAPPINGS = {
|
||||
"SRL Conditional Interrrupt": SrlConditionalInterrupt,
|
||||
"SRL Format String": SrlFormatString,
|
||||
"SRL Filter Image List": SrlFilterImageList,
|
||||
"SRL Num Expr": SrlNumExpr,
|
||||
"SRL Count SEGS": SrlCountSegs,
|
||||
"SRL Scale Image": SrlScaleImage,
|
||||
}
|
||||
|
||||
|
||||
@@ -159,5 +266,7 @@ NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"SrlConditionalInterrupt": "SRL Conditional Interrupt",
|
||||
"SrlFormatString": "SRL Format String",
|
||||
"SrlFilterImageList": "SRL Filter Image List",
|
||||
"SrlNumExpr": "SRL Num Expr",
|
||||
"SrlCountSegs": "SRL Count SEGS",
|
||||
"SrlScaleImage": "SRL Scale Image",
|
||||
}
|
||||
|
||||
+165
-9
@@ -20,50 +20,201 @@ def safe_getattr(obj: Any, attr: str) -> Any:
|
||||
if attr.startswith("_"):
|
||||
raise ValueError(f"Not allowed to access attributes starting with underscore.")
|
||||
|
||||
return getattr(obj, attr)
|
||||
value = getattr(obj, attr)
|
||||
if callable(value):
|
||||
if inspect.isbuiltin(value):
|
||||
if value.__qualname__ not in SAFE_BUILTIN_METHODS:
|
||||
raise ValueError(
|
||||
f"Sorry, {value.__qualname__!r} is not on the list of allowed builtin methods."
|
||||
)
|
||||
else:
|
||||
raise ValueError(f"Sorry, {value!r} is not on the list of allowed methods.")
|
||||
|
||||
return value
|
||||
|
||||
|
||||
SAFE_BUILTINS = {
|
||||
"abs": abs,
|
||||
"bytes": bytes,
|
||||
"all": all,
|
||||
"any": any,
|
||||
"ascii": ascii,
|
||||
"bin": bin,
|
||||
"bool": bool,
|
||||
"callable": callable,
|
||||
"chr": chr,
|
||||
"complex": complex,
|
||||
"dict": dict,
|
||||
"divmod": divmod,
|
||||
"format": format,
|
||||
"enumerate": enumerate,
|
||||
"filter": filter,
|
||||
"float": float,
|
||||
"format": format,
|
||||
"frozenset": frozenset,
|
||||
"getattr": safe_getattr,
|
||||
# TODO hasattr
|
||||
"hash": hash,
|
||||
"hex": hex,
|
||||
"id": id,
|
||||
"int": int,
|
||||
"isinstance": isinstance,
|
||||
"issubclass": issubclass,
|
||||
"iter": iter,
|
||||
"len": len,
|
||||
"list": list,
|
||||
"map": map,
|
||||
"max": max,
|
||||
"min": min,
|
||||
"next": next,
|
||||
"oct": oct,
|
||||
"ord": ord,
|
||||
"pow": pow,
|
||||
"range": range,
|
||||
"repr": repr,
|
||||
"reversed": reversed,
|
||||
"round": round,
|
||||
"set": set,
|
||||
"slice": slice,
|
||||
"sorted": sorted,
|
||||
"str": str,
|
||||
"sum": sum,
|
||||
"tuple": tuple,
|
||||
"zip": zip,
|
||||
}
|
||||
|
||||
|
||||
SAFE_BUILTIN_METHODS = set(
|
||||
[
|
||||
"complex.conjugate",
|
||||
"dict.iter",
|
||||
"dict.get",
|
||||
"dict.items",
|
||||
"dict.keys",
|
||||
"dict.reversed",
|
||||
"dict.values",
|
||||
"float.as_integer_ratio",
|
||||
"float.is_integer",
|
||||
"float.hex",
|
||||
"float.fromhex",
|
||||
"int.bit_length",
|
||||
"int.bit_count",
|
||||
"int.to_bytes",
|
||||
"int.from_bytes",
|
||||
"int.as_integer_ratio",
|
||||
"list.copy",
|
||||
"set.difference",
|
||||
"set.intersection",
|
||||
"set.isdisjoint",
|
||||
"set.issubset",
|
||||
"set.issuperset",
|
||||
"set.symmetric_difference",
|
||||
"set.union",
|
||||
"str.capitalize",
|
||||
"str.casefold",
|
||||
"str.center",
|
||||
"str.count",
|
||||
"str.encode",
|
||||
"str.endswith",
|
||||
"str.expandtabs",
|
||||
"str.find",
|
||||
"str.index",
|
||||
"str.isalnum",
|
||||
"str.isalpha",
|
||||
"str.isascii",
|
||||
"str.isdigit",
|
||||
"str.isidentifier",
|
||||
"str.islower",
|
||||
"str.isnumeric",
|
||||
"str.isprintable",
|
||||
"str.isspace",
|
||||
"str.istitle",
|
||||
"str.isupper",
|
||||
"str.join",
|
||||
"str.ljust",
|
||||
"str.lower",
|
||||
"str.lstrip",
|
||||
"str.maketrans",
|
||||
"str.partition",
|
||||
"str.removeprefix",
|
||||
"str.removesuffix",
|
||||
"str.replace",
|
||||
"str.rfind",
|
||||
"str.rindex",
|
||||
"str.rjust",
|
||||
"str.rpartition",
|
||||
"str.rstrip",
|
||||
"str.split",
|
||||
"str.splitlines",
|
||||
"str.startswith",
|
||||
"str.strip",
|
||||
"str.swapcase",
|
||||
"str.title",
|
||||
"str.translate",
|
||||
"str.upper",
|
||||
"str.zfill",
|
||||
]
|
||||
)
|
||||
|
||||
|
||||
class Safifier:
|
||||
"""Evaluate an expression in a restricted subset of Python by walking the AST."""
|
||||
|
||||
safe_nodes = (
|
||||
ast.Add,
|
||||
ast.And,
|
||||
ast.arg,
|
||||
ast.arguments,
|
||||
ast.Add,
|
||||
ast.BinOp,
|
||||
ast.BitAnd,
|
||||
ast.BitOr,
|
||||
ast.BitXor,
|
||||
ast.BoolOp,
|
||||
ast.Call,
|
||||
ast.Compare,
|
||||
ast.comprehension,
|
||||
ast.Constant,
|
||||
ast.Dict,
|
||||
ast.DictComp,
|
||||
ast.Div,
|
||||
ast.Eq,
|
||||
ast.Expression,
|
||||
ast.FloorDiv,
|
||||
ast.FormattedValue,
|
||||
ast.GeneratorExp,
|
||||
ast.Gt,
|
||||
ast.GtE,
|
||||
ast.IfExp,
|
||||
ast.In,
|
||||
ast.Invert,
|
||||
ast.Is,
|
||||
ast.IsNot,
|
||||
ast.JoinedStr,
|
||||
ast.keyword,
|
||||
ast.Lambda,
|
||||
ast.List,
|
||||
ast.ListComp,
|
||||
ast.Load,
|
||||
ast.LShift,
|
||||
ast.MatMult,
|
||||
ast.Mod,
|
||||
ast.Mult,
|
||||
ast.Not,
|
||||
ast.NotEq,
|
||||
ast.NotIn,
|
||||
ast.Name,
|
||||
ast.Or,
|
||||
ast.Pass,
|
||||
ast.Pow,
|
||||
ast.RShift,
|
||||
ast.Set,
|
||||
ast.SetComp,
|
||||
ast.Slice,
|
||||
ast.Starred,
|
||||
ast.Sub,
|
||||
ast.Subscript,
|
||||
ast.Tuple,
|
||||
ast.UAdd,
|
||||
ast.UnaryOp,
|
||||
ast.USub,
|
||||
)
|
||||
|
||||
def make_field_safe(self, field):
|
||||
@@ -96,8 +247,6 @@ class Safifier:
|
||||
else:
|
||||
raise NotImplementedError(f"Node {node!r} is not supported.")
|
||||
|
||||
ast.fix_missing_locations(new_node)
|
||||
|
||||
|
||||
def safe_compile(expr: str):
|
||||
node = ast.parse(expr, mode="eval")
|
||||
@@ -107,22 +256,29 @@ def safe_compile(expr: str):
|
||||
return compile(safe_node, "<string>", "eval")
|
||||
|
||||
|
||||
def evaluate(expr: str, local_symbols: dict[str, Any] = {}) -> Any:
|
||||
code = safe_compile(expr)
|
||||
def safe_eval(code, local_symbols: dict[str, Any] = {}) -> Any:
|
||||
global_symbols = {
|
||||
"__builtins__": SAFE_BUILTINS,
|
||||
}
|
||||
return eval(code, global_symbols, local_symbols)
|
||||
|
||||
|
||||
def evaluate(expr: str, local_symbols: dict[str, Any] = {}) -> Any:
|
||||
code = safe_compile(expr)
|
||||
return safe_eval(code, local_symbols)
|
||||
|
||||
|
||||
class SafeFormatter(string.Formatter):
|
||||
field_pat = re.compile(r"([^\.\]]*)\.(.+)|\[(\d+)\]")
|
||||
field_pat = re.compile(r"([^\.\]]*)(?:\.(.+)|\[(\d+)\])?")
|
||||
|
||||
def get_field(self, field_name, args, kwargs):
|
||||
m = self.field_pat.fullmatch(field_name)
|
||||
if m is None:
|
||||
raise ValueError(f"Could not parse field name {field_name!r}")
|
||||
|
||||
key, attr, idx = m.groups()
|
||||
if key[0].isdigit():
|
||||
key = int(key)
|
||||
value = self.get_value(key, args, kwargs)
|
||||
if attr is not None:
|
||||
return safe_getattr(value, attr), key
|
||||
|
||||
Reference in New Issue
Block a user