Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
6a88434dae | ||
|
|
074e59391b | ||
|
|
dfc34eb304 | ||
|
|
5ed48cfa02 | ||
|
|
ed442b59df | ||
|
|
3d6b2437d5 | ||
|
|
2514e340ef | ||
|
|
6082d2164e | ||
|
|
4df3c3e925 | ||
|
|
29698bb6db | ||
|
|
8bb4dd66c1 |
@@ -46,8 +46,9 @@ problems. PRs gladly accepted if you have need for this.
|
||||
|
||||

|
||||
|
||||
Takes a list of images and a list of bools as input and outputs a list
|
||||
of the input images where the bool is true.
|
||||
Takes a list of images and a list of bools as input and outputs two
|
||||
image lists, one where the bool is true and the other where it's
|
||||
false.
|
||||
|
||||
# License
|
||||
|
||||
|
||||
+146
-149
@@ -1,172 +1,169 @@
|
||||
import inspect
|
||||
import textwrap
|
||||
import itertools
|
||||
from typing import Annotated, Any, Literal, Optional
|
||||
|
||||
import comfy.utils
|
||||
import nodes
|
||||
|
||||
|
||||
class AnyType(str):
|
||||
def __ne__(self, __value: object) -> bool:
|
||||
return False
|
||||
from . import evaluator
|
||||
from .node_decorator import NODE_CLASS_MAPPINGS, NODE_NAME_MAPPINGS, Name, NumRange, node, IMAGE, SEGS
|
||||
|
||||
|
||||
any_typ = AnyType("*")
|
||||
|
||||
|
||||
class SrlConditionalInterrupt:
|
||||
@node("utils", "SRL Conditional Interrupt")
|
||||
def SrlConditionalInterrupt(interrupt: Annotated[bool, "forceInput"], inp: Any) -> Annotated[Any, Name("output")]:
|
||||
"""Interrupt processing if the boolean input is true. Pass through the other input."""
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"interrupt": ("BOOLEAN", {"forceInput": True}),
|
||||
"inp": (any_typ,),
|
||||
},
|
||||
}
|
||||
if interrupt:
|
||||
nodes.interrupt_processing()
|
||||
|
||||
RETURN_TYPES = (any_typ,)
|
||||
RETURN_NAMES = ("output",)
|
||||
FUNCTION = "doit"
|
||||
CATEGORY = "utils"
|
||||
|
||||
def doit(self, interrupt, inp):
|
||||
if interrupt:
|
||||
nodes.interrupt_processing()
|
||||
|
||||
return (inp,)
|
||||
return inp
|
||||
|
||||
|
||||
class SrlFormatString:
|
||||
"""Use Python f-string syntax to generate a string using the inputs as the arguments."""
|
||||
|
||||
@classmethod
|
||||
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}",
|
||||
}),
|
||||
},
|
||||
"optional": {
|
||||
"in0": (any_typ,),
|
||||
"in1": (any_typ,),
|
||||
"in2": (any_typ,),
|
||||
"in3": (any_typ,),
|
||||
"in4": (any_typ,),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("STRING",)
|
||||
FUNCTION = "doit"
|
||||
CATEGORY = "utils"
|
||||
|
||||
def doit(self, format, **kwargs):
|
||||
# Allow referencing arguments both by name and index.
|
||||
return (format.format(*kwargs.values(), **kwargs),)
|
||||
@node("utils", "SRL Format String")
|
||||
def SrlFormatString(
|
||||
format: str = "first input via str(): {}, second input via repr(): {!r}, third input by index: {2}",
|
||||
in0: Optional[Any] = None,
|
||||
in1: Optional[Any] = None,
|
||||
in2: Optional[Any] = None,
|
||||
in3: Optional[Any] = None,
|
||||
in4: Optional[Any] = None,
|
||||
) -> str:
|
||||
return evaluator.safe_format(format, in0, in1, in2, in3, in4)
|
||||
|
||||
|
||||
class SrlEval:
|
||||
"""Evaluate any Python code as a function with the given inputs."""
|
||||
@node("utils", "SRL Num Expr")
|
||||
class SrlNumExpr:
|
||||
"""Evaluate a numerical expression safely."""
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"parameters": ("STRING", {
|
||||
"multiline": False,
|
||||
"default": "a, b=None, c=\"foo\", *rest",
|
||||
"dynamicPrompts": False,
|
||||
}),
|
||||
"code": ("STRING", {
|
||||
"multiline": True,
|
||||
"default": "code goes here\nreturn a + b",
|
||||
"dynamicPrompts": False,
|
||||
}),
|
||||
},
|
||||
"optional": {
|
||||
"arg0": (any_typ,),
|
||||
"arg1": (any_typ,),
|
||||
"arg2": (any_typ,),
|
||||
"arg3": (any_typ,),
|
||||
"arg4": (any_typ,),
|
||||
}
|
||||
}
|
||||
def __init__(self):
|
||||
self.last_expr = None
|
||||
self.last_code = None
|
||||
|
||||
RETURN_TYPES = (any_typ,)
|
||||
FUNCTION = "doit"
|
||||
CATEGORY = "utils"
|
||||
def __call__(
|
||||
self,
|
||||
expr: Annotated[str, "multiline"],
|
||||
i0: Optional[Annotated[int, "forceInput"]] = None,
|
||||
i1: Optional[Annotated[int, "forceInput"]] = None,
|
||||
i2: Optional[Annotated[int, "forceInput"]] = None,
|
||||
f0: Optional[Annotated[float, "forceInput"]] = None,
|
||||
f1: Optional[Annotated[float, "forceInput"]] = None,
|
||||
f2: Optional[Annotated[float, "forceInput"]] = None,
|
||||
) -> tuple[bool, int, float]:
|
||||
if expr == self.last_expr:
|
||||
code = self.last_code
|
||||
else:
|
||||
code = self.last_code = evaluator.safe_compile(expr)
|
||||
|
||||
def doit(self, parameters, code, **kw):
|
||||
# Indent the code for the main body of the function
|
||||
func_code = textwrap.indent(code, " ")
|
||||
source = f"def func({parameters}):\n{func_code}"
|
||||
|
||||
# The provided code can mutate globals or really do anything, but ComfyUI isn't secure to begin with.
|
||||
loc = {}
|
||||
exec(source, globals(), loc)
|
||||
func = loc["func"]
|
||||
|
||||
argspec = inspect.getfullargspec(func)
|
||||
# We don't allow variable keyword arguments or keyword only arguments, but we do allow varargs
|
||||
assert argspec.varkw is None
|
||||
assert not argspec.kwonlyargs
|
||||
|
||||
input_names = list(self.INPUT_TYPES()["optional"].keys())
|
||||
parameter_names = argspec.args
|
||||
|
||||
# Convert the list of defaults into a dictionary to make it easier to use
|
||||
default_list = argspec.defaults if argspec.defaults is not None else []
|
||||
defaults = {parameter_name: default for parameter_name, default in zip(parameter_names[-len(default_list):], default_list)}
|
||||
|
||||
# We handle substituting default values ourselves in order to support *args
|
||||
args = [kw[input_name] if input_name in kw else defaults[parameter_name] for parameter_name, input_name in zip(parameter_names, input_names)]
|
||||
|
||||
# Support *args
|
||||
if argspec.varargs is not None:
|
||||
unnamed_inputs = input_names[len(argspec.args):]
|
||||
# I considered requiring the remaining inputs to be contiguous, but I don't think it's helpful.
|
||||
args += [kw[input_name] for input_name in unnamed_inputs if input_name in kw]
|
||||
|
||||
ret = func(*args)
|
||||
return (ret,)
|
||||
res = evaluator.safe_eval(code, {
|
||||
"i0": i0,
|
||||
"i1": i1,
|
||||
"i2": i2,
|
||||
"f0": f0,
|
||||
"f1": f1,
|
||||
"f2": f2,
|
||||
})
|
||||
return (bool(res), int(res), float(res))
|
||||
|
||||
|
||||
class SrlFilterImageList:
|
||||
@node("utils", "SRL Filter Image List")
|
||||
def SrlFilterImageList(images: list[IMAGE], keep: list[Annotated[bool, "forceInput"]]) -> tuple[list[Annotated[IMAGE, Name("t_images")]], list[Annotated[IMAGE, Name("f_images")]]]:
|
||||
"""Filter an image list based on a list of bools"""
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"images": ("IMAGE",),
|
||||
"keep": ("BOOLEAN", {"forceInput": True}),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("IMAGE",)
|
||||
INPUT_IS_LIST = True
|
||||
OUTPUT_IS_LIST = (True,)
|
||||
FUNCTION = "doit"
|
||||
|
||||
def doit(self, images, keep):
|
||||
return ([im for im, k in zip(images, keep) if k],)
|
||||
im1, im2 = itertools.tee(zip(images, keep))
|
||||
return (
|
||||
[im for im, k in im1 if k],
|
||||
[im for im, k in im2 if not k],
|
||||
)
|
||||
|
||||
|
||||
# A dictionary that contains all nodes you want to export with their names
|
||||
# NOTE: names should be globally unique
|
||||
NODE_CLASS_MAPPINGS = {
|
||||
"SRL Conditional Interrrupt": SrlConditionalInterrupt,
|
||||
"SRL Format String": SrlFormatString,
|
||||
"SRL Eval": SrlEval,
|
||||
"SRL Filter Image List": SrlFilterImageList,
|
||||
}
|
||||
@node("utils", "SRL Count SEGS")
|
||||
def SrlCountSegs(segs: SEGS, ignore_none: bool = False) -> int:
|
||||
"""Count the number of segs."""
|
||||
if segs is None:
|
||||
if ignore_none:
|
||||
raise TypeError("segs is None, expected SEGS")
|
||||
else:
|
||||
return 0
|
||||
else:
|
||||
return len(segs[1])
|
||||
|
||||
|
||||
# A dictionary that contains the friendly/humanly readable titles for the nodes
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"SrlConditionalInterrupt": "SRL Conditional Interrupt",
|
||||
"SrlFormatString": "SRL Format String",
|
||||
"SrlEval": "SRL Eval",
|
||||
"SrlFilterImageList": "SRL Filter Image List",
|
||||
}
|
||||
@node("images", "SRL Scale Image")
|
||||
class SrlScaleImage:
|
||||
"""Scale an image under specified conditions."""
|
||||
|
||||
def scale_dim(self, how, idim, odim, width, height) -> tuple[int, int, float, bool, bool]:
|
||||
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 round(width), round(height), factor, scaled, scaled_up
|
||||
|
||||
def __call__(
|
||||
self,
|
||||
image: IMAGE,
|
||||
what: Literal["width", "height", "both", "factor", "longest", "shortest", "mebipixels"],
|
||||
how: Literal["up", "down", "both"],
|
||||
width: Annotated[int, NumRange(1, nodes.MAX_RESOLUTION, 1)],
|
||||
height: Annotated[int, NumRange(1, nodes.MAX_RESOLUTION, 1)],
|
||||
factor: Annotated[float, NumRange(0.001, 100.0, 0.001)],
|
||||
longest_shortest: Annotated[int, NumRange(1, nodes.MAX_RESOLUTION, 1)],
|
||||
mebipixels: Annotated[float, NumRange(0.01, 20, 0.01)],
|
||||
) -> tuple[IMAGE, Annotated[int, Name("width")], Annotated[int, Name("height")], Annotated[float, Name("width_factor")], Annotated[float, Name("height_factor")], Annotated[bool, Name("did_scale")], Annotated[bool, Name("did_scale_up")]]:
|
||||
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
|
||||
case _:
|
||||
raise ValueError("'what' must be width, height, longest, shortest, or mebipixels")
|
||||
|
||||
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)
|
||||
|
||||
+306
@@ -0,0 +1,306 @@
|
||||
"""Evaluate Python expressions safely."""
|
||||
|
||||
import ast
|
||||
import inspect
|
||||
from operator import getitem
|
||||
import re
|
||||
import string
|
||||
from typing import Any, Generator, Optional, Sequence
|
||||
|
||||
|
||||
def safe_vformat(fmt: str, args: Sequence = [], kwargs: dict = {}) -> str:
|
||||
return FORMATTER.vformat(fmt, args, kwargs)
|
||||
|
||||
|
||||
def safe_format(fmt: str, *args, **kwargs) -> str:
|
||||
return safe_vformat(fmt, args, kwargs)
|
||||
|
||||
|
||||
def safe_getattr(obj: Any, attr: str) -> Any:
|
||||
if attr.startswith("_"):
|
||||
raise ValueError(f"Not allowed to access attributes starting with underscore.")
|
||||
|
||||
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,
|
||||
"all": all,
|
||||
"any": any,
|
||||
"ascii": ascii,
|
||||
"bin": bin,
|
||||
"bool": bool,
|
||||
"callable": callable,
|
||||
"chr": chr,
|
||||
"complex": complex,
|
||||
"dict": dict,
|
||||
"divmod": divmod,
|
||||
"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.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):
|
||||
if field is None or isinstance(field, (int, str)):
|
||||
return field
|
||||
|
||||
if isinstance(field, ast.AST):
|
||||
return self.make_safe(field)
|
||||
|
||||
if isinstance(field, list):
|
||||
return [self.make_safe(node) for node in field]
|
||||
|
||||
raise NotImplementedError(f"Field {field!r} not supported.")
|
||||
|
||||
def make_safe(self, node: ast.AST) -> ast.AST:
|
||||
if isinstance(node, ast.Attribute):
|
||||
assert isinstance(node.ctx, ast.Load)
|
||||
value = self.make_safe(node.value)
|
||||
return ast.Call(
|
||||
func=ast.Name(id="getattr", ctx=ast.Load()),
|
||||
args=[
|
||||
value,
|
||||
ast.Constant(value=node.attr),
|
||||
],
|
||||
keywords=[],
|
||||
)
|
||||
elif isinstance(node, self.safe_nodes):
|
||||
fields = [self.make_field_safe(f) for _, f in ast.iter_fields(node)]
|
||||
return node.__class__(*fields)
|
||||
else:
|
||||
raise NotImplementedError(f"Node {node!r} is not supported.")
|
||||
|
||||
|
||||
def safe_compile(expr: str):
|
||||
node = ast.parse(expr, mode="eval")
|
||||
safifier = Safifier()
|
||||
safe_node = safifier.make_safe(node)
|
||||
ast.fix_missing_locations(safe_node)
|
||||
return compile(safe_node, "<string>", "eval")
|
||||
|
||||
|
||||
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+)\])?")
|
||||
|
||||
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
|
||||
|
||||
if idx is not None:
|
||||
return getitem(value, int(idx)), key
|
||||
|
||||
return value, key
|
||||
|
||||
|
||||
FORMATTER = SafeFormatter()
|
||||
|
||||
|
||||
def main():
|
||||
from argparse import ArgumentParser
|
||||
|
||||
p = ArgumentParser()
|
||||
p.add_argument("expression")
|
||||
args = p.parse_args()
|
||||
result = evaluate(args.expression, {})
|
||||
print(f"result: {result!r}")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,228 @@
|
||||
from dataclasses import dataclass, field
|
||||
import inspect
|
||||
import types
|
||||
from typing import Any, Annotated, Literal, Optional, TypeAlias, Union, get_args, get_origin
|
||||
|
||||
|
||||
NODE_CLASS_MAPPINGS = {}
|
||||
NODE_NAME_MAPPINGS = {}
|
||||
|
||||
|
||||
@dataclass
|
||||
class NumRange:
|
||||
min: int | float
|
||||
max: int | float
|
||||
step: int | float
|
||||
|
||||
|
||||
@dataclass
|
||||
class Name:
|
||||
name: str
|
||||
|
||||
|
||||
@dataclass
|
||||
class TypeName:
|
||||
"""Annotation for a wrapper around arbitrary types like 'IMAGE'"""
|
||||
name: str
|
||||
|
||||
|
||||
class AnyType(str):
|
||||
def __ne__(self, __value: object) -> bool:
|
||||
return False
|
||||
|
||||
|
||||
IMAGE: TypeAlias = Annotated[Any, TypeName("IMAGE")]
|
||||
SEGS: TypeAlias = Annotated[Any, TypeName("SEGS")]
|
||||
|
||||
|
||||
def process_annotation(a) -> tuple[str | tuple, Optional[str], bool, bool, dict[str, Any]]:
|
||||
settings = {}
|
||||
type_name = None
|
||||
name = None
|
||||
is_list = False
|
||||
optional = False
|
||||
|
||||
print(f"Processing annotation for {a!r}")
|
||||
|
||||
if get_origin(a) in (Union, types.UnionType):
|
||||
args = get_args(a)
|
||||
if len(args) != 2 or args[1] is not type(None):
|
||||
raise TypeError(f"We only support unions with NoneType, i.e. Optional")
|
||||
|
||||
a = args[0]
|
||||
optional = True
|
||||
|
||||
if get_origin(a) is list:
|
||||
is_list = True
|
||||
a = get_args(a)[0]
|
||||
|
||||
if get_origin(a) is Annotated:
|
||||
for arg in get_args(a)[1:]:
|
||||
print(f"arg: {arg!r}")
|
||||
match arg:
|
||||
case "dynamicPrompts":
|
||||
settings["dynamicPrompts"] = True
|
||||
case "forceInput":
|
||||
settings["forceInput"] = True
|
||||
case "multiline":
|
||||
settings["multiline"] = True
|
||||
case NumRange(min, max, step):
|
||||
settings.update({
|
||||
"min": min,
|
||||
"max": max,
|
||||
"step": step,
|
||||
})
|
||||
case Name(n):
|
||||
name = n
|
||||
case TypeName(n):
|
||||
type_name = n
|
||||
|
||||
a = get_args(a)[0]
|
||||
|
||||
if type_name is None:
|
||||
print(f"Checking type {a!r}")
|
||||
# TODO validate settings
|
||||
if a is int:
|
||||
type_name = "INT"
|
||||
elif a is bool:
|
||||
type_name = "BOOLEAN"
|
||||
elif a is float:
|
||||
type_name = "FLOAT"
|
||||
elif a is str:
|
||||
type_name = "STRING"
|
||||
settings.setdefault("multiline", False)
|
||||
settings.setdefault("dynamicPrompts", False)
|
||||
elif a is Any:
|
||||
type_name = AnyType("*")
|
||||
elif get_origin(a) is Literal:
|
||||
type_name = get_args(a)
|
||||
else:
|
||||
raise ValueError(f"Need to provide TypeName annotation for {a!r}, or you annotated the outer type")
|
||||
|
||||
return type_name, name, optional, is_list, settings
|
||||
|
||||
|
||||
def process_return(a) -> tuple[str, str, bool]:
|
||||
t, name, optional, is_list, settings = process_annotation(a)
|
||||
|
||||
# No literal returns
|
||||
assert not isinstance(t, tuple)
|
||||
|
||||
if optional:
|
||||
raise TypeError(f"Optional makes no sense for return types in {t!r} {name!r}")
|
||||
|
||||
if name is None:
|
||||
name = t
|
||||
|
||||
return t, name, is_list
|
||||
|
||||
|
||||
def node(category, name=None, input_is_list=None):
|
||||
def wrapper(func):
|
||||
nonlocal input_is_list, name
|
||||
tuplize = False
|
||||
|
||||
if isinstance(func, type):
|
||||
# Support callable classes
|
||||
sig = inspect.signature(func.__call__)
|
||||
d = func.__dict__.copy()
|
||||
bases = func.__bases__ # Keep any existing bases
|
||||
skip_self = True
|
||||
func = func.__call__
|
||||
def doit1(self, *args, **kwargs):
|
||||
r = func(self, *args, **kwargs)
|
||||
return (r,) if tuplize else r
|
||||
d.update({
|
||||
"FUNCTION": "doit",
|
||||
"doit": doit1,
|
||||
})
|
||||
else:
|
||||
def doit2(self, *args, **kwargs):
|
||||
r = func(*args, **kwargs)
|
||||
return (r,) if tuplize else r
|
||||
|
||||
sig = inspect.signature(func)
|
||||
d = {
|
||||
"FUNCTION": "doit",
|
||||
"doit": doit2,
|
||||
}
|
||||
bases = ()
|
||||
skip_self = False
|
||||
|
||||
arg_is_list = False
|
||||
required = {}
|
||||
optional = {}
|
||||
return_types = []
|
||||
return_names = []
|
||||
output_is_list = []
|
||||
|
||||
ra = sig.return_annotation
|
||||
if get_origin(ra) is not tuple:
|
||||
# Turn singletons into a tuple to simplify processing
|
||||
ra = tuple[ra]
|
||||
tuplize = True
|
||||
elif get_origin(ra) is not tuple:
|
||||
raise TypeError("Return type of node must be a type or tuple of types.")
|
||||
|
||||
for rt in get_args(ra):
|
||||
return_type, return_name, is_list = process_return(rt)
|
||||
return_types.append(return_type)
|
||||
return_names.append(return_name)
|
||||
output_is_list.append(is_list)
|
||||
|
||||
for arg_name, parameter in list(sig.parameters.items())[1 if skip_self else 0:]:
|
||||
print(f"arg_name = {arg_name!r}, parameter = {parameter!r}")
|
||||
type_name, n, opt, is_list, settings = process_annotation(parameter.annotation)
|
||||
if n:
|
||||
raise TypeError(f"Don't use Name annotation on input types: {arg_name!r}")
|
||||
|
||||
if is_list:
|
||||
if input_is_list is None:
|
||||
input_is_list = True
|
||||
arg_is_list = True
|
||||
elif not input_is_list:
|
||||
raise TypeError(f"All inputs must be is_list or none of them, or use second argument of the decorator: {arg_name!r}")
|
||||
elif arg_is_list:
|
||||
raise TypeError(f"All inputs must be is_list or none of them, or use second argument of the decorator: {arg_name!r}")
|
||||
|
||||
if parameter.default is not inspect.Parameter.empty and parameter.default is not None:
|
||||
settings["default"] = parameter.default
|
||||
|
||||
if isinstance(type_name, tuple):
|
||||
if settings:
|
||||
raise TypeError("No settings with list of selections: {n!r}")
|
||||
o = type_name
|
||||
elif settings:
|
||||
o = (type_name, settings)
|
||||
else:
|
||||
o = (type_name,)
|
||||
|
||||
if opt:
|
||||
optional[arg_name] = o
|
||||
else:
|
||||
required[arg_name] = o
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": required,
|
||||
"optional": optional,
|
||||
}
|
||||
|
||||
d.update({
|
||||
"CATEGORY": category,
|
||||
"INPUT_TYPES": INPUT_TYPES,
|
||||
"INPUT_IS_LIST": False if input_is_list is None else input_is_list,
|
||||
"RETURN_TYPES": return_types,
|
||||
"RETURN_NAMES": return_names,
|
||||
"OUTPUT_IS_LIST": output_is_list,
|
||||
})
|
||||
new_node = type(func.__name__, bases, d)
|
||||
if name is None:
|
||||
name = new_node.__name__
|
||||
|
||||
NODE_CLASS_MAPPINGS[name] = new_node
|
||||
NODE_NAME_MAPPINGS[new_node.__name__] = name
|
||||
return new_node
|
||||
|
||||
return wrapper
|
||||
@@ -0,0 +1,76 @@
|
||||
import unittest
|
||||
from typing import Annotated, Any, Literal, Optional
|
||||
|
||||
from node_decorator import NODE_CLASS_MAPPINGS, NODE_NAME_MAPPINGS, Name, NumRange, node, IMAGE, SEGS
|
||||
|
||||
|
||||
@node("utils", "SRL Conditional Interrupt")
|
||||
def SrlConditionalInterrupt(interrupt: Annotated[bool, "forceInput"], inp: Any) -> Annotated[Any, Name("output")]:
|
||||
...
|
||||
|
||||
|
||||
@node("utils", "SRL Format String")
|
||||
def SrlFormatString(
|
||||
format: str = "first input via str(): {}, second input via repr(): {!r}, third input by index: {2}",
|
||||
in0: Optional[Any] = None,
|
||||
in1: Optional[Any] = None,
|
||||
in2: Optional[Any] = None,
|
||||
in3: Optional[Any] = None,
|
||||
in4: Optional[Any] = None,
|
||||
) -> str:
|
||||
...
|
||||
|
||||
|
||||
@node("utils", "SRL Num Expr")
|
||||
class SrlNumExpr:
|
||||
def __call__(
|
||||
self,
|
||||
expr: Annotated[str, "multiline"],
|
||||
i0: Optional[Annotated[int, "forceInput"]],
|
||||
i1: Optional[Annotated[int, "forceInput"]],
|
||||
i2: Optional[Annotated[int, "forceInput"]],
|
||||
f0: Optional[Annotated[float, "forceInput"]],
|
||||
f1: Optional[Annotated[float, "forceInput"]],
|
||||
f2: Optional[Annotated[float, "forceInput"]]
|
||||
) -> tuple[bool, int, float]:
|
||||
...
|
||||
|
||||
|
||||
@node("utils", "SRL Filter Image List")
|
||||
def SrlFilterImageList(images: list[IMAGE], keep: list[Annotated[bool, "forceInput"]]) -> tuple[list[Annotated[IMAGE, Name("t_images")]], list[Annotated[IMAGE, Name("f_images")]]]:
|
||||
...
|
||||
|
||||
|
||||
@node("utils", "SRL Count SEGS")
|
||||
def SrlCountSegs(segs: SEGS, ignore_none: bool = False) -> int:
|
||||
...
|
||||
|
||||
|
||||
@node("images", "SRL Scale Image")
|
||||
class SrlScaleImage:
|
||||
def __call__(
|
||||
self,
|
||||
image: IMAGE,
|
||||
what: Literal["width", "height", "both", "factor", "longest", "shortest", "mebipixels"],
|
||||
how: Literal["up", "down", "both"],
|
||||
width: Annotated[int, NumRange(1, 4096, 1)],
|
||||
height: Annotated[int, NumRange(1, 4096, 1)],
|
||||
factor: Annotated[float, NumRange(0.001, 100.0, 0.001)],
|
||||
longest_shortest: Annotated[int, NumRange(1, 4096, 1)],
|
||||
mebipixels: Annotated[float, NumRange(0.01, 20, 0.01)],
|
||||
) -> tuple[IMAGE, Annotated[int, Name("width")], Annotated[int, Name("height")], Annotated[float, Name("width_factor")], Annotated[float, Name("height_factor")], Annotated[bool, Name("did_scale")], Annotated[bool, Name("did_scale_up")]]:
|
||||
...
|
||||
|
||||
|
||||
def main():
|
||||
print(f"NODE_CLASS_MAPPINGS = {NODE_CLASS_MAPPINGS!r}")
|
||||
print(f"NODE_NAME_MAPPINGS = {NODE_NAME_MAPPINGS!r}")
|
||||
|
||||
for name, cls in NODE_CLASS_MAPPINGS.items():
|
||||
print(name)
|
||||
print(f" {cls.INPUT_TYPES()!r}")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
|
||||
@@ -0,0 +1,35 @@
|
||||
import unittest
|
||||
|
||||
from evaluator import evaluate, safe_vformat
|
||||
|
||||
|
||||
class TestEvaluator(unittest.TestCase):
|
||||
def test_fstring_access(self):
|
||||
with self.assertRaises(ValueError):
|
||||
evaluate("(lambda x: f'{x.__class__}')(1)")
|
||||
|
||||
def test_fstrings(self):
|
||||
s = {"x": 5}
|
||||
r = evaluate("f'{x}'", s)
|
||||
self.assertEqual(r, "5")
|
||||
r = evaluate("f'{x:05}'", s)
|
||||
self.assertEqual(r, "00005")
|
||||
|
||||
def test_lambda(self):
|
||||
r = evaluate("(lambda x, y: x + y)(2, 3)")
|
||||
self.assertEqual(r, 5)
|
||||
|
||||
def test_attribute_access(self):
|
||||
with self.assertRaises(ValueError):
|
||||
evaluate("str.__class__")
|
||||
|
||||
|
||||
class TestFormatter(unittest.TestCase):
|
||||
def test_format_access(self):
|
||||
with self.assertRaises(ValueError):
|
||||
s = safe_vformat("{x.__class__}", [], {"x": object})
|
||||
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
Reference in New Issue
Block a user