From ed442b59df080ae46787598d5a9d1631c1e94ac5 Mon Sep 17 00:00:00 2001 From: Sean Lynch Date: Mon, 2 Oct 2023 19:33:55 -0400 Subject: [PATCH] Latest changes --- __init__.py | 135 +++++++++++++++++++++++++++++++++++---- evaluator.py | 174 ++++++++++++++++++++++++++++++++++++++++++++++++--- 2 files changed, 287 insertions(+), 22 deletions(-) diff --git a/__init__.py b/__init__.py index ce67a2e..bcf4a1a 100644 --- a/__init__.py +++ b/__init__.py @@ -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", } diff --git a/evaluator.py b/evaluator.py index 978b681..e1f5fef 100644 --- a/evaluator.py +++ b/evaluator.py @@ -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, "", "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