Latest changes

This commit is contained in:
Sean Lynch
2023-10-02 19:33:55 -04:00
parent 3d6b2437d5
commit ed442b59df
2 changed files with 287 additions and 22 deletions
+122 -13
View File
@@ -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
View File
@@ -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