11 Commits
Author SHA1 Message Date
Sean Lynch 6a88434dae Fix decorated classes
Also add defaults to SrlNumExpr
2023-10-22 18:35:20 -04:00
Sean Lynch 074e59391b Make optional actually work 2023-10-16 21:04:01 -04:00
Sean Lynch dfc34eb304 Fix STRING
Added multiline and set default settings for strings.
2023-10-16 20:59:24 -04:00
Sean Lynch 5ed48cfa02 Node decorator 2023-10-16 20:48:23 -04:00
Sean Lynch ed442b59df Latest changes 2023-10-02 19:33:55 -04:00
Sean Lynch 3d6b2437d5 Mutate the AST and compile instead
As long as we only let nodes through that we believe to be safe, it
should be fine to let Python handle the execution. The main unsafe
piece of code is attribute lookup.

Right now this allows access to str.format. Need to fix that.
2023-09-30 18:26:45 -04:00
Sean Lynch 2514e340ef Last commit before attempting compiler 2023-09-30 17:19:56 -04:00
Sean Lynch 6082d2164e Implement a safe(r) string formatter 2023-09-30 15:02:44 -04:00
Sean Lynch 4df3c3e925 Add tests, improve implementation
Dramatically improve the readability of the implementation by not
trying to use the visitor from ast.
2023-09-30 13:16:35 -04:00
Sean Lynch 29698bb6db Properly handle slices missing start or end 2023-09-29 17:51:17 -04:00
Sean Lynch 8bb4dd66c1 Initial implementation of Python expression evaluator
I haven't tested it particularly heavily yet. Need to find or write a
Python expression conformance suite.
2023-09-29 17:45:29 -04:00
8 changed files with 794 additions and 186 deletions
-21
View File
@@ -1,21 +0,0 @@
name: Publish to Comfy registry
on:
workflow_dispatch:
push:
branches:
- main
paths:
- "pyproject.toml"
jobs:
publish-node:
name: Publish Custom Node to registry
runs-on: ubuntu-latest
steps:
- name: Check out code
uses: actions/checkout@v4
- name: Publish Custom Node
uses: Comfy-Org/publish-node-action@main
with:
## Add your own personal access token to your Github Repository secrets and reference it here.
personal_access_token: ${{ secrets.REGISTRY_ACCESS_TOKEN }}
+3 -2
View File
@@ -46,8 +46,9 @@ problems. PRs gladly accepted if you have need for this.
![Screenshot of SrlFilterImageList](screenshots/SrlFilterImageList.png)
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
View File
@@ -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
View File
@@ -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()
+228
View File
@@ -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
-14
View File
@@ -1,14 +0,0 @@
[project]
name = "srl-nodes"
description = "Nodes: SRL Conditional Interrupt, SRL Format String, SRL Eval, SRL Filter Image List. This is a collection of nodes I find useful. Note that at least one module allows execution of arbitrary code. Do not use any of these nodes on a system that allow untrusted users to control workflows or inputs.[w/WARNING: The custom nodes in this extension are vulnerable to **security risks** because they allow the execution of arbitrary code through the workflow]"
version = "1.0.0"
license = "AGPLv3"
[project.urls]
Repository = "https://github.com/seanlynch/srl-nodes"
# Used by Comfy Registry https://comfyregistry.org
[tool.comfy]
PublisherId = "seanlynch"
DisplayName = "srl-nodes"
Icon = ""
+76
View File
@@ -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()
+35
View File
@@ -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()