Compare commits

...
41 Commits
Author SHA1 Message Date
Mel Massadian 2441f19db3 ⬆️ Bump version: 0.2.0 → 0.2.1 2024-12-30 18:46:35 +01:00
Mel Massadian b01e027ec0 Merge branch 'main' into pr-227 2024-12-30 18:40:13 +01:00
Mel Massadian 9a4ecb2b90 fix: 🐛 handle missing submodules
the nodes should never fail to load completely.
I still need to remove the few remaining side effects like this one.
2024-12-30 18:37:43 +01:00
Mel Massadian 0eeb707f34 feat: ✨ add SaveImage passthrough
Exactly like the native one but not as an OUTPUT_NODE,
primarly meant to "inline" image saving in upcoming mtb loops.
2024-12-30 18:29:28 +01:00
Mel Massadian 9a943714aa chore: 🧹 dev
dev files
2024-12-29 13:51:03 +01:00
Robin Huang 885688e7c7 Checkout submodules before publishing. 2024-12-27 14:39:48 -08:00
Mel Massadian bae26a07fb feat: ✨ add filtering to TransformImage
fixes #209
2024-12-27 19:29:34 +01:00
Mel Massadian c92d99a8a3 feat: ✨ add support for video in I/O sidebar
slow if you have big videos, maybe it shouldn't use force_size
2024-12-22 05:43:19 +01:00
Mel Massadian 3f6d082940 feat: ✨ add an extra static input to Stack Images 2024-12-22 02:40:01 +01:00
Mel Massadian 58ae89f8e0 chore: 🧹 apply formatting 2024-12-22 02:40:01 +01:00
pak c9a26427a8 improve dynamic inputs: custom separator and start_index, preserve labels, ... 2024-12-22 02:40:01 +01:00
Mel Massadian 6608c0b6d1 fix: 🐛 add warnings about what each IO mode can do
VHS now has a Load Image Path node that could be used to solve all cases.
For this I'll need to get the full path of each images from the endpoint
2024-12-22 02:11:21 +01:00
Mel Massadian a757e1c98b fix: 🐛 soft deprecate compression h264 2024-12-22 01:30:03 +01:00
Mel Massadian 52bd76e19c feat: ✨ add support for subdirs (i/o sidebar)
fixes #221
2024-12-22 01:28:08 +01:00
Mel Massadian d6e004cce2 fix: 🐛 limit packages allowed to be installed from API
fixes #224

thanks @boy-hack for the report!
2024-12-22 00:22:24 +01:00
Mel Massadian ed17fa2ef4 fix: 🐛 ensure default settings (io sidebar)
fixes #225
2024-12-21 02:14:40 +01:00
Mel Massadian 827c64c43d feat: ✨ add Batch Sequence Nodes
- A regular one that just sequence batches
- A "plus" with transition support (POC + for now)
2024-12-16 01:44:01 +01:00
filtered e5482aee5e fix: 🐛 spawn colour picker at pointer location (#223) 2024-12-15 22:22:22 +01:00
Mel Massadian 62469a4dd9 fix: 🐛 i/o sidebar for custom paths
In utils I uses a constant for these which doesn't
update with the global... calling the getters should
solve that.

This issue is probably in other places where I use these
utils.

fixes #219
2024-12-11 00:09:43 +01:00
Mel Massadian 8c629bee18 feat: ✨ add support for more formats (I/O sidebar) 2024-12-08 23:13:07 +01:00
Mel Massadian 50cb6f5ed6 chore: 🧹 bump minor 2024-12-08 19:34:26 +01:00
Mel Massadian e32d1e02df feat: ✨ add h264 compression node
recommended for i2i in ltx.
original code by [mix](https://github.com/XmYx)
2024-12-08 19:12:28 +01:00
Mel Massadian b0d52f7305 fix: 🐛 remove mtb sidebar
- The source for this is not yet in main... this file slipped
  in an earlier commit

fixes https://github.com/Comfy-Org/ComfyUI_frontend/issues/1834
2024-12-07 15:45:43 +01:00
Mel Massadian e17c6e29f5 docs: 📚 update wiki
pull wiki for documentation
2024-12-04 02:11:02 +01:00
Mel Massadian 27e03fa23e feat: ✨ add postshot nodes
basic wrapper of the cli the idea is to
queue Cog + Rotating loras -> Postshot

needs testing
2024-12-03 23:17:54 +01:00
Mel Massadian ec1cb1ac17 fix: 🐛 always enable the I/O sidebar
closes #214
2024-12-03 22:11:48 +01:00
d8ahazard 64634104a2 Use local import for Rembg
Rembg can sometimes cause *very* long load times on import (like 40s). Moving it to local doesn't fix the long import entirely, but it does prevent it causing ComfyUI from loading slowly.
2024-12-03 04:55:55 +01:00
Mel Massadian ecbb220de6 fix: 🐛 ui shifts on animation builder
finally updated to addDOMWidget
2024-11-20 23:03:00 +01:00
Mel Massadian cd9e614b1a feat: ✨ improve the I/O sidebar
- better options (sort, count)
- uses the new toast api instead of MTB.notify
2024-11-20 22:42:32 +01:00
Mel Massadian 9ccf572a15 chore: 🧹 add worktree to gitignores
for the experimental doc site at:
https://melmass.github.io/comfy_mtb/
2024-11-20 22:42:32 +01:00
Mel Massadian 74af5c6499 feat: ✨ add UpscaleBBoxBy 2024-11-20 22:42:32 +01:00
Mel Massadian caf0b39d8a chore 🧹: add deprecations and experimental 2024-11-20 22:42:32 +01:00
Mel Massadian e099d581a7 chore: 🧹 remove dupe code 2024-11-20 22:42:32 +01:00
Mel Massadian 22f7c30373 feat: ✨ simplified sidebar and backend
If you have a LoadImage selected,
clicking on images in the "input" mode will set the image on the
selected nodes
2024-11-20 22:42:32 +01:00
Mel Massadian 0133fb93bc feat: ✨ add Interpolate Condition 2024-11-20 22:42:32 +01:00
Mel Massadian cf7d30507e feat: ✨ dump of wip things... 2024-11-20 22:42:32 +01:00
Mel Massadian b6fa571fd2 fix: 🐛 category for settings 2024-11-20 21:41:57 +01:00
Mel Massadian f272526bfc fix: 🐛 new UI issues
- Fixes the "edit icon cannot be clicked"
- Changed the parser to add support for more non std markdown
- Markdown links now always open a new tab instead of replacing current
- New optional shiki support for code blocks (check #211 for details)
2024-11-20 21:41:57 +01:00
Mel Massadian 4e593bb30b feat: ✨ use the new parser for documentations
- might also fix #210
2024-11-20 21:41:57 +01:00
Mel Massadian 097ca33b8e feat: ✨ add @mtb/markdown-parser bundles
- the standard one is half the size of showdown
- the enhanced one (add shiki with most of its features) is 1.5mb
2024-11-20 21:41:57 +01:00
Mel Massadian 784fb0145b chore: 🧹 update externs
- remove showdown
- update dompurify
2024-11-20 21:41:57 +01:00
36 changed files with 4246 additions and 1169 deletions
+2
View File
@@ -12,6 +12,8 @@ jobs:
steps: steps:
- name: ♻️ Check out code - name: ♻️ Check out code
uses: actions/checkout@v4 uses: actions/checkout@v4
with:
submodules: true
- name: 📦 Publish Custom Node - name: 📦 Publish Custom Node
uses: Comfy-Org/publish-node-action@main uses: Comfy-Org/publish-node-action@main
with: with:
+3
View File
@@ -6,3 +6,6 @@ node_modules/
compose.yaml compose.yaml
comfy_mtb.wsb comfy_mtb.wsb
Dockerfile Dockerfile
# I store the gh-pages worktrees (src & build) there
.worktrees
+209 -57
View File
@@ -7,10 +7,12 @@
# #
### ###
__version__ = "0.1.6" __version__ = "0.2.1"
import os import os
from aiohttp.web_request import Request
# TODO: don't override this if the user has that setup already # TODO: don't override this if the user has that setup already
if not os.environ.get("TF_FORCE_GPU_ALLOW_GROWTH"): if not os.environ.get("TF_FORCE_GPU_ALLOW_GROWTH"):
os.environ["TF_FORCE_GPU_ALLOW_GROWTH"] = "true" os.environ["TF_FORCE_GPU_ALLOW_GROWTH"] = "true"
@@ -31,24 +33,21 @@ from pathlib import Path
from aiohttp import web from aiohttp import web
from server import PromptServer from server import PromptServer
import nodes
from .endpoint import endlog from .endpoint import endlog
from .install import get_node_dependencies
from .log import blue_text, cyan_text, get_label, get_summary, log from .log import blue_text, cyan_text, get_label, get_summary, log
from .utils import comfy_dir, here from .utils import comfy_dir, here
NODE_CLASS_MAPPINGS = {} NODE_CLASS_MAPPINGS: dict[str, type] = {}
NODE_DISPLAY_NAME_MAPPINGS = {} NODE_DISPLAY_NAME_MAPPINGS: dict[str, str] = {}
NODE_CLASS_MAPPINGS_DEBUG = {} NODE_CLASS_MAPPINGS_DEBUG: dict[str, str | None] = {}
WEB_DIRECTORY = "./web" WEB_DIRECTORY = "./web"
def extract_nodes_from_source(filename: Path): def extract_nodes_from_source(filename: Path):
source_code = "" source_code = ""
source_code = filename.read_text(encoding="utf-8") source_code = filename.read_text(encoding="utf-8")
nodes: list[str] = []
nodes = []
try: try:
parsed = ast.parse(source_code) parsed = ast.parse(source_code)
@@ -57,14 +56,15 @@ def extract_nodes_from_source(filename: Path):
target = node.targets[0] target = node.targets[0]
if isinstance(target, ast.Name) and target.id == "__nodes__": if isinstance(target, ast.Name) and target.id == "__nodes__":
value = ast.get_source_segment(source_code, node.value) value = ast.get_source_segment(source_code, node.value)
node_value = ast.parse(value).body[0].value if value:
if isinstance(node_value, (ast.List, ast.Tuple)): node_value = ast.parse(value).body[0].value
nodes.extend( if isinstance(node_value, ast.List | ast.Tuple):
element.id nodes.extend(
for element in node_value.elts str(element.id)
if isinstance(element, ast.Name) for element in node_value.elts
) if isinstance(element, ast.Name)
break )
break
except SyntaxError: except SyntaxError:
log.error("Failed to parse") log.error("Failed to parse")
return nodes return nodes
@@ -72,8 +72,8 @@ def extract_nodes_from_source(filename: Path):
def load_nodes(): def load_nodes():
errors: list[str] = [] errors: list[str] = []
nodes = [] nodes: list[type] = []
nodes_failed = [] nodes_failed: list[str] = []
for filename in (here / "nodes").iterdir(): for filename in (here / "nodes").iterdir():
if filename.suffix == ".py": if filename.suffix == ".py":
@@ -124,7 +124,8 @@ def uninstall_old_web_extensions():
shutil.rmtree(web_mtb) shutil.rmtree(web_mtb)
except Exception as e: except Exception as e:
log.warning( log.warning(
f"Failed to remove web mtb directory: {e}\nPlease manually remove it from disk ({web_mtb}) and restart the server." f"""Failed to remove web mtb directory: {e}
Please manually remove it from disk ({web_mtb}) and restart the server."""
) )
@@ -141,7 +142,7 @@ def wiki_to_classname(s: str):
def classname_to_wiki(s: str): def classname_to_wiki(s: str):
classname = s.replace("MTB_", "") classname = s.replace("MTB_", "")
parts = [] parts: list[str] = []
start = 0 start = 0
for i in range(1, len(classname)): for i in range(1, len(classname)):
if classname[i].isupper(): if classname[i].isupper():
@@ -161,8 +162,6 @@ if wiki.exists() and wiki.is_dir():
# - REGISTER NODES # - REGISTER NODES
MTB_EXPORT = os.environ.get("MTB_EXPORT") MTB_EXPORT = os.environ.get("MTB_EXPORT")
nodes, failed = load_nodes() nodes, failed = load_nodes()
@@ -179,7 +178,7 @@ for node_class in nodes:
node_class.DESCRIPTION = node_class.__doc__ node_class.DESCRIPTION = node_class.__doc__
if MTB_EXPORT: if MTB_EXPORT:
wiki_name = classname_to_wiki(class_name) wiki_name = classname_to_wiki(class_name)
(wiki / "nodes" / (wiki_name + ".md")).write_text( _ = (wiki / "nodes" / (wiki_name + ".md")).write_text(
node_class.__doc__, encoding="utf-8" node_class.__doc__, encoding="utf-8"
) )
@@ -192,12 +191,15 @@ for node_class in nodes:
NODE_CLASS_MAPPINGS[node_label] = node_class NODE_CLASS_MAPPINGS[node_label] = node_class
NODE_DISPLAY_NAME_MAPPINGS[class_name] = node_label NODE_DISPLAY_NAME_MAPPINGS[class_name] = node_label
NODE_CLASS_MAPPINGS_DEBUG[node_label] = node_class.__doc__ NODE_CLASS_MAPPINGS_DEBUG[node_label] = node_class.__doc__
# TODO: I removed this, I find it more convenient to write without spaces, but it breaks every of my workflows
# TODO (cont): and until I find a way to automate the conversion, I'll leave it like this # TODO: I removed this, I find it more convenient to write without spaces
# but it breaks every of my workflows
# TODO (cont): and until I find a way to automate the conversion
# I'll leave it like this
if os.environ.get("MTB_EXPORT"): if os.environ.get("MTB_EXPORT"):
with open(here / "node_list.json", "w") as f: with open(here / "node_list.json", "w") as f:
f.write( _ = f.write(
json.dumps( json.dumps(
{ {
k: NODE_CLASS_MAPPINGS_DEBUG[k] k: NODE_CLASS_MAPPINGS_DEBUG[k]
@@ -215,29 +217,31 @@ log.debug(
) )
) )
log.info(f"loaded {cyan_text(len(nodes))} nodes successfuly") log.info(f"loaded {cyan_text(str(len(nodes)))} nodes successfuly")
if failed: if failed:
with contextlib.suppress(Exception): with contextlib.suppress(Exception):
base_url, port = utils.get_server_info() base_url, port = utils.get_server_info()
log.info( log.info(
f"Some nodes ({len(failed)}) could not be loaded. This can be ignored, but go to http://{base_url}:{port}/mtb if you want more information." f"Some nodes ({len(failed)}) could not be loaded. This can be ignored, but go to http://{base_url}:{port}/mtb if you want more information."
) )
log.debug(failed)
# - ENDPOINT # - ENDPOINT
if hasattr(PromptServer, "instance"): if hasattr(PromptServer, "instance"):
restore_deps = ["basicsr"] img_cache = None
onnx_deps = ["onnxruntime"] prompt_cache = None
swap_deps = ["insightface"] + onnx_deps
node_dependency_mapping = { with contextlib.suppress(ImportError):
"QrCode": ["qrcode"], from cachetools import TTLCache
"DeepBump": onnx_deps,
"FaceSwap": swap_deps, img_cache = TTLCache(maxsize=100, ttl=5) # 1 min TTL
"LoadFaceSwapModel": swap_deps, prompt_cache = TTLCache(maxsize=100, ttl=5) # 1 min TTL
"LoadFaceAnalysisModel": restore_deps,
} node_dependency_mapping = get_node_dependencies()
PromptServer.instance.app.router.add_static( PromptServer.instance.app.router.add_static(
"/mtb-assets/", path=(here / "html").as_posix() "/mtb-assets/", path=(here / "html").as_posix()
@@ -306,10 +310,10 @@ if hasattr(PromptServer, "instance"):
} }
) )
@PromptServer.instance.routes.post("/mtb/debug") @PromptServer.instance.routes.post("/mtb/server-info")
async def set_debug(request): async def set_server_info(request: Request):
json_data = await request.json() json_data: dict[str, bool] = await request.json()
enabled = json_data.get("enabled") enabled = json_data.get("debug")
if enabled: if enabled:
os.environ["MTB_DEBUG"] = "true" os.environ["MTB_DEBUG"] = "true"
log.setLevel(logging.DEBUG) log.setLevel(logging.DEBUG)
@@ -317,7 +321,7 @@ if hasattr(PromptServer, "instance"):
elif "MTB_DEBUG" in os.environ: elif "MTB_DEBUG" in os.environ:
# del os.environ["MTB_DEBUG"] # del os.environ["MTB_DEBUG"]
os.environ.pop("MTB_DEBUG") _ = os.environ.pop("MTB_DEBUG")
log.setLevel(logging.INFO) log.setLevel(logging.INFO)
return web.json_response( return web.json_response(
@@ -325,17 +329,17 @@ if hasattr(PromptServer, "instance"):
) )
@PromptServer.instance.routes.get("/mtb") @PromptServer.instance.routes.get("/mtb")
async def get_home(request): async def get_home(request: Request):
from . import endpoint from . import endpoint
reload(endpoint) _ = reload(endpoint)
# Check if the request prefers HTML content # Check if the request prefers HTML content
if "text/html" in request.headers.get("Accept", ""): if "text/html" in request.headers.get("Accept", ""):
# # Return an HTML page # # Return an HTML page
html_response = """ html_response = """
<div class="flex-container menu"> <div class="flex-container menu">
<a href="/mtb/manage">manage</a> <a href="/mtb/manage">manage</a>
<a href="/mtb/debug">debug</a> <a href="/mtb/server-info">Server Info</a>
<a href="/mtb/status">status</a> <a href="/mtb/status">status</a>
</div> </div>
""" """
@@ -347,28 +351,176 @@ if hasattr(PromptServer, "instance"):
# Return JSON for other requests # Return JSON for other requests
return web.json_response({"message": "Welcome to MTB!"}) return web.json_response({"message": "Welcome to MTB!"})
@PromptServer.instance.routes.get("/mtb/debug") import asyncio
async def get_debug(request): import os
from io import BytesIO
from aiohttp import web
from PIL import Image
def get_cached_image(file_path: str, preview_params=None, channel=None):
cache_key = (file_path, preview_params, channel)
if img_cache and (cache_key in img_cache):
return img_cache[cache_key]
with Image.open(file_path) as img:
info = img.info
if preview_params:
img = process_preview(img, preview_params)
if channel:
img = process_channel(img, channel)
if prompt_cache:
prompt_cache[cache_key] = info
if img_cache:
img_cache[cache_key] = img.getvalue()
return img_cache[cache_key]
return img.getvalue()
def process_preview(img: Image.Image, preview_params):
image_format, quality, width = preview_params
quality = int(quality)
if width:
width = int(width)
img.thumbnail((width, int(width * img.height / img.width)))
buffer = BytesIO()
img.save(
buffer, format=image_format, quality=quality, metadata=img.info
)
buffer.seek(0)
return buffer
def process_channel(img: Image.Image, channel: str):
if channel == "rgb":
if img.mode == "RGBA":
r, g, b, _ = img.split()
img = Image.merge("RGB", (r, g, b))
else:
img = img.convert("RGB")
elif channel == "a":
if img.mode == "RGBA":
_, _, _, a = img.split()
else:
a = Image.new("L", img.size, 255)
img = Image.new("RGBA", img.size)
img.putalpha(a)
buffer = BytesIO()
img.save(buffer, format="PNG")
_ = buffer.seek(0)
return buffer
async def get_image_response(
file, filename: str, preview_info=None, channel=None
):
img = await asyncio.to_thread(
get_cached_image, file, preview_info, channel
)
return web.Response(
body=img,
content_type="image/webp" if preview_info else "image/png",
headers={"Content-Disposition": f'filename="{filename}"'},
)
# TODO: Embed the metadatas somehow so we can drag and drop
# to load workflows in the sidebar
@PromptServer.instance.routes.get("/mtb/view")
async def view_image(request: Request):
import folder_paths
filename = request.rel_url.query.get("filename")
if not filename:
return web.Response(status=404)
filename, output_dir = folder_paths.annotated_filepath(filename)
if filename[0] == "/" or ".." in filename:
return web.Response(status=400)
if output_dir is None:
rtype = request.rel_url.query.get("type", "output")
output_dir = folder_paths.get_directory_by_type(rtype)
if output_dir is None:
return web.Response(status=400)
if "subfolder" in request.rel_url.query:
full_output_dir = os.path.join(
output_dir, request.rel_url.query["subfolder"]
)
if (
os.path.commonpath(
(os.path.abspath(full_output_dir), output_dir)
)
!= output_dir
):
return web.Response(status=403)
output_dir = full_output_dir
filename = os.path.basename(filename)
file = os.path.join(output_dir, filename)
if not os.path.isfile(file):
return web.Response(status=404)
preview_info = None
if "preview" in request.rel_url.query:
preview_params = request.rel_url.query["preview"].split(";")
image_format = (
preview_params[0]
if preview_params[0] in ["webp", "jpeg"]
else "webp"
)
quality = (
int(preview_params[1])
if len(preview_params) > 1 and preview_params[1].isdigit()
else 90
)
width = request.rel_url.query.get("width")
preview_info = (image_format, quality, width)
channel = request.rel_url.query.get("channel")
return await get_image_response(file, filename, preview_info, channel)
@PromptServer.instance.routes.get("/mtb/server-info")
async def get_debug(request: Request):
from . import endpoint from . import endpoint
reload(endpoint) _ = reload(endpoint)
enabled = "MTB_DEBUG" in os.environ isdebug = "MTB_DEBUG" in os.environ
exposed = "MTB_EXPOSE" in os.environ
def render_property(name: str, val: str):
return f"""<strong>{name}:</strong>
<p>
{val}
</p>"""
# Check if the request prefers HTML content # Check if the request prefers HTML content
if "text/html" in request.headers.get("Accept", ""): if "text/html" in request.headers.get("Accept", ""):
# # Return an HTML page # # Return an HTML page
html_response = f""" html_response = ""
<h1>MTB Debug Status: {'Enabled' if enabled else 'Disabled'}</h1>
""" html_response += render_property(
"Debug", "Enabled" if isdebug else "Disabled"
)
html_response += render_property("Exposed", str(exposed))
return web.Response( return web.Response(
text=endpoint.render_base_template("Debug", html_response), text=endpoint.render_base_template(
"Server Info", html_response
),
content_type="text/html", content_type="text/html",
) )
# Return JSON for other requests # Return JSON for other requests
return web.json_response({"enabled": enabled}) return web.json_response({"exposed": exposed, "debug": isdebug})
@PromptServer.instance.routes.get("/mtb/actions") @PromptServer.instance.routes.get("/mtb/actions")
async def no_route(request): async def no_route(request: Request):
from . import endpoint from . import endpoint
if "text/html" in request.headers.get("Accept", ""): if "text/html" in request.headers.get("Accept", ""):
@@ -382,7 +534,7 @@ if hasattr(PromptServer, "instance"):
return web.json_response({"message": "actions has no get for now"}) return web.json_response({"message": "actions has no get for now"})
@PromptServer.instance.routes.post("/mtb/actions") @PromptServer.instance.routes.post("/mtb/actions")
async def do_action(request): async def do_action(request: Request):
from . import endpoint from . import endpoint
reload(endpoint) reload(endpoint)
+28 -19
View File
@@ -1,22 +1,31 @@
{ {
"$schema": "https://biomejs.dev/schemas/1.6.1/schema.json", "$schema": "https://biomejs.dev/schemas/1.6.1/schema.json",
"organizeImports": { "organizeImports": {
"enabled": true "enabled": true
}, },
"linter": { "linter": {
"enabled": true, "enabled": true,
"rules": { "rules": {
"recommended": true "recommended": true,
} "suspicious": {
}, "noConsoleLog": "warn"
"formatter": { },
"lineEnding": "lf" "style": {
}, "noParameterAssign": "off",
"javascript": { "noShoutyConstants": "warn",
"formatter": { "useNamingConvention": "off"
"quoteStyle": "single", }
"semicolons": "asNeeded",
"indentWidth": 2
}
} }
},
"formatter": {
"indentStyle": "space",
"indentWidth": 2,
"lineEnding": "lf"
},
"javascript": {
"formatter": {
"quoteStyle": "single",
"semicolons": "asNeeded"
}
}
} }
+152 -20
View File
@@ -1,10 +1,20 @@
import csv import csv
import secrets
import sys
import urllib.parse
from pathlib import Path
from typing import Any, Literal
import folder_paths
from aiohttp import web from aiohttp import web
from .install import get_node_dependencies
from .log import mklog from .log import mklog
from .utils import ( from .utils import (
SortMode,
backup_file, backup_file,
build_glob_patterns,
glob_multiple,
import_install, import_install,
reqs_map, reqs_map,
run_command, run_command,
@@ -14,18 +24,25 @@ from .utils import (
endlog = mklog("mtb endpoint") endlog = mklog("mtb endpoint")
# - ACTIONS # - ACTIONS
import sys
from pathlib import Path
import_install("requirements") import_install("requirements")
def ACTIONS_installDependency(dependency_names=None): def ACTIONS_installDependency(dependency_names: list[str] | None = None):
if dependency_names is None: if dependency_names is None:
# return web.Response(text="No dependency name provided", status=400)
return {"error": "No dependency name provided"} return {"error": "No dependency name provided"}
endlog.debug(f"Received Install Dependency request for {dependency_names}") endlog.debug(f"Received Install Dependency request for {dependency_names}")
# reqs = [] # reqs = []
resolved_names = [reqs_map.get(name, name) for name in dependency_names] resolved_names = [reqs_map.get(name, name) for name in dependency_names]
allowed_deps = list(
{d for dep in get_node_dependencies().values() for d in dep}
)
for dep in dependency_names:
if dep not in allowed_deps:
return {
"error": f"Unknown dependency: {dep}, you can only use this endpoint to install {allowed_deps}"
}
try: try:
run_command( run_command(
[Path(sys.executable), "-m", "pip", "install"] + resolved_names [Path(sys.executable), "-m", "pip", "install"] + resolved_names
@@ -50,6 +67,106 @@ def ACTIONS_installDependency(dependency_names=None):
# break # break
def ACTIONS_getUserImageFolders():
input_dir = Path(folder_paths.get_input_directory())
output_dir = Path(folder_paths.get_output_directory())
input_subdirs = [x.name for x in input_dir.iterdir() if x.is_dir()]
output_subdirs = [x.name for x in output_dir.iterdir() if x.is_dir()]
return {"input": input_subdirs, "output": output_subdirs}
def ACTIONS_getUserVideos(
size=256, count=200, offset=0, sort: str | None = None
):
count = count or 1000
video_extensions = ["webm", "mp4", "mkv", "mov"]
entries = {}
patterns = build_glob_patterns(video_extensions)
input_dir = Path(folder_paths.get_input_directory())
entries = glob_multiple(input_dir, patterns)
sort_mode = SortMode.from_str(sort)
if sort_mode:
sort_key = {
SortMode.MODIFIED: lambda x: x.stat().st_mtime,
SortMode.MODIFIED_REVERSE: lambda x: x.stat().st_mtime,
SortMode.NAME: lambda x: x.name,
SortMode.NAME_REVERSE: lambda x: x.name,
}.get(sort_mode)
if sort_key:
reverse = sort_mode in (SortMode.MODIFIED, SortMode.NAME_REVERSE)
entries = sorted(entries, key=sort_key, reverse=reverse)
videos = {
video.name: (
f"/view?force_rate=0&frame_load_cap=0&skip_first_frames=0&select_every_nth=1&filename={urllib.parse.quote_plus(video.name)}&type=input&format=video&force_size={size}x?"
)
for i, video in enumerate(entries)
if offset <= i < offset + count
}
return videos
def ACTIONS_getUserImages(
mode: Literal["input", "output"],
count=1000,
offset=0,
sort: str | None = None,
include_subfolders: bool = False,
subfolder=None,
):
# enabled = "MTB_EXPOSE" in os.environ
# if not enabled:
# return {"error": "Session not authorized to getInputs"}
imgs = {}
count = count or 1000
input_dir = Path(folder_paths.get_input_directory())
output_dir = Path(folder_paths.get_output_directory())
entry_dir = input_dir if mode == "input" else output_dir
if subfolder:
entry_dir = entry_dir / subfolder
if not entry_dir.exists():
return {
"error": f"Subfolder {entry_dir.name} doesn't exists in {entry_dir.parent.as_posix()}"
}
supported = ["png", "jpg", "jpeg", "webp", "gif"]
entries = {}
patterns = build_glob_patterns(supported, recursive=include_subfolders)
entries = glob_multiple(entry_dir, patterns)
sort_mode = SortMode.from_str(sort)
if sort_mode:
sort_key = {
SortMode.MODIFIED: lambda x: x.stat().st_mtime,
SortMode.MODIFIED_REVERSE: lambda x: x.stat().st_mtime,
SortMode.NAME: lambda x: x.name,
SortMode.NAME_REVERSE: lambda x: x.name,
}.get(sort_mode)
if sort_key:
reverse = sort_mode in (SortMode.MODIFIED, SortMode.NAME_REVERSE)
entries = sorted(entries, key=sort_key, reverse=reverse)
imgs = {
img.name: (
f"/mtb/view?filename={img.name}&width=512&type={mode}&subfolder={subfolder or ''}"
f"{img.parent.relative_to(entry_dir) if include_subfolders else ''}"
f"&preview=&rand={secrets.randbelow(424242)}"
)
for i, img in enumerate(entries)
if offset <= i < offset + count
}
return imgs
def ACTIONS_getStyles(style_name=None): def ACTIONS_getStyles(style_name=None):
from .nodes.conditions import MTB_StylesLoader from .nodes.conditions import MTB_StylesLoader
@@ -97,7 +214,7 @@ def ACTIONS_saveStyle(data):
csv_writer.writerow(row) csv_writer.writerow(row)
async def do_action(request) -> web.Response: async def do_action(request: web.Request) -> web.Response:
endlog.debug("Init action request") endlog.debug("Init action request")
request_data = await request.json() request_data = await request.json()
name = request_data.get("name") name = request_data.get("name")
@@ -109,7 +226,12 @@ async def do_action(request) -> web.Response:
method = globals().get(method_name) method = globals().get(method_name)
if callable(method): if callable(method):
result = method(args) if args else method() result = None
if args:
result = method(*args) if isinstance(args, list) else method(args)
else:
result = method()
endlog.debug(f"Action result: {result}") endlog.debug(f"Action result: {result}")
return web.json_response({"result": result}) return web.json_response({"result": result})
@@ -130,10 +252,13 @@ async def do_action(request) -> web.Response:
# - HTML UTILS # - HTML UTILS
def dependencies_button(name, dependencies): def dependencies_button(name: str, dependencies: list[str]) -> str:
deps = ",".join([f"'{x}'" for x in dependencies]) deps = ",".join([f"'{x}'" for x in dependencies])
return f""" return f"""
<button class="dependency-button" onclick="window.mtb_action('installDependency',[{deps}])">Install {name} deps</button> <button
class="dependency-button"
onclick="window.mtb_action('installDependency',[{deps}])"
>Install {name} deps</button>
""" """
@@ -153,7 +278,7 @@ def csv_editor():
html_out = """ html_out = """
<div id="style-editor"> <div id="style-editor">
<h1>Style Editor</h1> <h1>Style Editor</h1>
""" """
for current, styles in style_files.items(): for current, styles in style_files.items():
current_out = f"<h3>{current}</h3>" current_out = f"<h3>{current}</h3>"
@@ -215,11 +340,14 @@ def render_tab_view(**kwargs):
""" """
def add_foldable_region(title, content): def add_foldable_region(title: str, content: str):
symbol_id = f"{title}-symbol" symbol_id = f"{title}-symbol"
return f""" return f"""
<div class='foldable'> <div class='foldable'>
<div class='foldable-title' onclick="toggleFoldable('{title}', '{symbol_id}')"> <div
class='foldable-title'
onclick="toggleFoldable('{title}', '{symbol_id}')"
>
<span id='{symbol_id}' class='foldable-symbol'>&#9655;</span> <span id='{symbol_id}' class='foldable-symbol'>&#9655;</span>
{title} {title}
</div> </div>
@@ -231,7 +359,9 @@ def add_foldable_region(title, content):
""" """
def add_split_pane(left_content, right_content, vertical=True): def add_split_pane(
left_content: str, right_content: str, *, vertical: bool = True
):
orientation = "vertical" if vertical else "horizontal" orientation = "vertical" if vertical else "horizontal"
return f""" return f"""
<div class="split-pane {orientation}"> <div class="split-pane {orientation}">
@@ -250,7 +380,7 @@ def add_split_pane(left_content, right_content, vertical=True):
""" """
def add_dropdown(title, options): def add_dropdown(title: str, options: list[str]):
option_str = "\n".join( option_str = "\n".join(
[f"<option value='{opt}'>{opt}</option>" for opt in options] [f"<option value='{opt}'>{opt}</option>" for opt in options]
) )
@@ -262,13 +392,13 @@ def add_dropdown(title, options):
""" """
def render_table(table_dict, sort=True, title=None): def render_table(table_dict: dict[str, Any], sort=True, title=None):
table_dict = sorted( table_list = sorted(
table_dict.items(), key=lambda item: item[0] table_dict.items(), key=lambda item: item[0]
) # Sort the dictionary by keys ) # Sort the dictionary by keys
table_rows = "" table_rows = ""
for name, item in table_dict: for name, item in table_list:
if isinstance(item, dict): if isinstance(item, dict):
if "dependencies" in item: if "dependencies" in item:
table_rows += f"<tr><td>{name}</td><td>" table_rows += f"<tr><td>{name}</td><td>"
@@ -299,12 +429,12 @@ def render_table(table_dict, sort=True, title=None):
<tbody> <tbody>
{table_rows} {table_rows}
</tbody> </tbody>
</table> </table>
</div> </div>
""" """
def render_base_template(title, content): def render_base_template(title: str, content: str):
github_icon_svg = """<svg xmlns="http://www.w3.org/2000/svg" fill="whitesmoke" height="3em" viewBox="0 0 496 512"><path d="M165.9 397.4c0 2-2.3 3.6-5.2 3.6-3.3.3-5.6-1.3-5.6-3.6 0-2 2.3-3.6 5.2-3.6 3-.3 5.6 1.3 5.6 3.6zm-31.1-4.5c-.7 2 1.3 4.3 4.3 4.9 2.6 1 5.6 0 6.2-2s-1.3-4.3-4.3-5.2c-2.6-.7-5.5.3-6.2 2.3zm44.2-1.7c-2.9.7-4.9 2.6-4.6 4.9.3 2 2.9 3.3 5.9 2.6 2.9-.7 4.9-2.6 4.6-4.6-.3-1.9-3-3.2-5.9-2.9zM244.8 8C106.1 8 0 113.3 0 252c0 110.9 69.8 205.8 169.5 239.2 12.8 2.3 17.3-5.6 17.3-12.1 0-6.2-.3-40.4-.3-61.4 0 0-70 15-84.7-29.8 0 0-11.4-29.1-27.8-36.6 0 0-22.9-15.7 1.6-15.4 0 0 24.9 2 38.6 25.8 21.9 38.6 58.6 27.5 72.9 20.9 2.3-16 8.8-27.1 16-33.7-55.9-6.2-112.3-14.3-112.3-110.5 0-27.5 7.6-41.3 23.6-58.9-2.6-6.5-11.1-33.3 2.6-67.9 20.9-6.5 69 27 69 27 20-5.6 41.5-8.5 62.8-8.5s42.8 2.9 62.8 8.5c0 0 48.1-33.6 69-27 13.7 34.7 5.2 61.4 2.6 67.9 16 17.7 25.8 31.5 25.8 58.9 0 96.5-58.9 104.2-114.8 110.5 9.2 7.9 17 22.9 17 46.4 0 33.7-.3 75.4-.3 83.6 0 6.5 4.6 14.4 17.3 12.1C428.2 457.8 496 362.9 496 252 496 113.3 383.5 8 244.8 8zM97.2 352.9c-1.3 1-1 3.3.7 5.2 1.6 1.6 3.9 2.3 5.2 1 1.3-1 1-3.3-.7-5.2-1.6-1.6-3.9-2.3-5.2-1zm-10.8-8.1c-.7 1.3.3 2.9 2.3 3.9 1.6 1 3.6.7 4.3-.7.7-1.3-.3-2.9-2.3-3.9-2-.6-3.6-.3-4.3.7zm32.4 35.6c-1.6 1.3-1 4.3 1.3 6.2 2.3 2.3 5.2 2.6 6.5 1 1.3-1.3.7-4.3-1.3-6.2-2.2-2.3-5.2-2.6-6.5-1zm-11.4-14.7c-1.6 1-1.6 3.6 0 5.9 1.6 2.3 4.3 3.3 5.6 2.3 1.6-1.3 1.6-3.9 0-6.2-1.4-2.3-4-3.3-5.6-2z"/></svg>""" github_icon_svg = """<svg xmlns="http://www.w3.org/2000/svg" fill="whitesmoke" height="3em" viewBox="0 0 496 512"><path d="M165.9 397.4c0 2-2.3 3.6-5.2 3.6-3.3.3-5.6-1.3-5.6-3.6 0-2 2.3-3.6 5.2-3.6 3-.3 5.6 1.3 5.6 3.6zm-31.1-4.5c-.7 2 1.3 4.3 4.3 4.9 2.6 1 5.6 0 6.2-2s-1.3-4.3-4.3-5.2c-2.6-.7-5.5.3-6.2 2.3zm44.2-1.7c-2.9.7-4.9 2.6-4.6 4.9.3 2 2.9 3.3 5.9 2.6 2.9-.7 4.9-2.6 4.6-4.6-.3-1.9-3-3.2-5.9-2.9zM244.8 8C106.1 8 0 113.3 0 252c0 110.9 69.8 205.8 169.5 239.2 12.8 2.3 17.3-5.6 17.3-12.1 0-6.2-.3-40.4-.3-61.4 0 0-70 15-84.7-29.8 0 0-11.4-29.1-27.8-36.6 0 0-22.9-15.7 1.6-15.4 0 0 24.9 2 38.6 25.8 21.9 38.6 58.6 27.5 72.9 20.9 2.3-16 8.8-27.1 16-33.7-55.9-6.2-112.3-14.3-112.3-110.5 0-27.5 7.6-41.3 23.6-58.9-2.6-6.5-11.1-33.3 2.6-67.9 20.9-6.5 69 27 69 27 20-5.6 41.5-8.5 62.8-8.5s42.8 2.9 62.8 8.5c0 0 48.1-33.6 69-27 13.7 34.7 5.2 61.4 2.6 67.9 16 17.7 25.8 31.5 25.8 58.9 0 96.5-58.9 104.2-114.8 110.5 9.2 7.9 17 22.9 17 46.4 0 33.7-.3 75.4-.3 83.6 0 6.5 4.6 14.4 17.3 12.1C428.2 457.8 496 362.9 496 252 496 113.3 383.5 8 244.8 8zM97.2 352.9c-1.3 1-1 3.3.7 5.2 1.6 1.6 3.9 2.3 5.2 1 1.3-1 1-3.3-.7-5.2-1.6-1.6-3.9-2.3-5.2-1zm-10.8-8.1c-.7 1.3.3 2.9 2.3 3.9 1.6 1 3.6.7 4.3-.7.7-1.3-.3-2.9-2.3-3.9-2-.6-3.6-.3-4.3.7zm32.4 35.6c-1.6 1.3-1 4.3 1.3 6.2 2.3 2.3 5.2 2.6 6.5 1 1.3-1.3.7-4.3-1.3-6.2-2.2-2.3-5.2-2.6-6.5-1zm-11.4-14.7c-1.6 1-1.6 3.6 0 5.9 1.6 2.3 4.3 3.3 5.6 2.3 1.6-1.3 1.6-3.9 0-6.2-1.4-2.3-4-3.3-5.6-2z"/></svg>"""
return f""" return f"""
<!DOCTYPE html> <!DOCTYPE html>
@@ -340,7 +470,9 @@ def render_base_template(title, content):
<header> <header>
<a href="/">Back to Comfy</a> <a href="/">Back to Comfy</a>
<div class="mtb_logo"> <div class="mtb_logo">
<img src="https://repository-images.githubusercontent.com/649047066/a3eef9a7-20dd-4ef9-b839-884502d4e873" alt="Comfy MTB Logo" height="70" width="128"> <img
src="https://repository-images.githubusercontent.com/649047066/a3eef9a7-20dd-4ef9-b839-884502d4e873"
alt="Comfy MTB Logo" height="70" width="128">
<span class="title">Comfy MTB</span></div> <span class="title">Comfy MTB</span></div>
<a style="width:128px;text-align:center" href="https://www.github.com/melmass/comfy_mtb"> <a style="width:128px;text-align:center" href="https://www.github.com/melmass/comfy_mtb">
{github_icon_svg} {github_icon_svg}
@@ -355,6 +487,6 @@ def render_base_template(title, content):
<!-- Shared footer content here --> <!-- Shared footer content here -->
</footer> </footer>
</body> </body>
</html> </html>
""" """
+17 -5
View File
@@ -27,7 +27,7 @@ export def "comfy start" [--clean,--old-ui, --listen] {
let root = get_root --clean=($clean) let root = get_root --clean=($clean)
cd $root cd $root
MTB_DEBUG=true python main.py --port 3000 ...(if $old_ui { ["--front-end-version", "Comfy-Org/ComfyUI_legacy_frontend@latest"]} else {[]}) --preview-method auto ...(if $listen {["--listen"]} else {[]}) MTB_DEBUG=true python main.py --port 3000 ...(if $old_ui { ["--front-end-version", "Comfy-Org/ComfyUI_legacy_frontend@latest"]} else {[ --front-end-version Comfy-Org/ComfyUI_frontend@latest]}) --preview-method auto ...(if $listen {["--listen"]} else {[]})
} }
# update comfy itself and merge master in current branch # update comfy itself and merge master in current branch
@@ -67,8 +67,14 @@ export def "comfy update" [
git checkout master git checkout master
print $"(ansi yellow_italic)Fetching and pulling remote updates(ansi reset)" print $"(ansi yellow_italic)Fetching and pulling remote updates(ansi reset)"
git fetch if ($clean) {
git pull git fetch local master
git pull local master
} else {
git fetch
git pull
}
print $"(ansi yellow_italic)Back to our branch \(($branch_name)\)(ansi reset)" print $"(ansi yellow_italic)Back to our branch \(($branch_name)\)(ansi reset)"
git checkout - git checkout -
@@ -135,7 +141,7 @@ export def "comfy update_extensions" [--clean] {
let root = get_root --clean=($clean) let root = get_root --clean=($clean)
cd $root cd $root
cd custom_nodes cd custom_nodes
git multipull . git multipull . -s -q
} }
def --env path-add [pth] { def --env path-add [pth] {
@@ -146,7 +152,7 @@ def --env path-add [pth] {
export-env { export-env {
$env.COMFY_MTB = ("." | path expand) $env.COMFY_MTB = ("." | path expand)
$env.CUDA_ROOT = 'C:\Program Files\NVIDIA GPU Computing Toolkit\CUDA\v12.1\' # $env.CUDA_ROOT = 'C:\Program Files\NVIDIA GPU Computing Toolkit\CUDA\v12.1\'
$env.CUDA_HOME = $env.CUDA_ROOT $env.CUDA_HOME = $env.CUDA_ROOT
@@ -154,6 +160,12 @@ export-env {
$env.COMFY_CLEAN_ROOT = ($env.COMFY_ROOT | path dirname | path join ComfyClean) $env.COMFY_CLEAN_ROOT = ($env.COMFY_ROOT | path dirname | path join ComfyClean)
path-add 'C:/Portable/TensorRT-8.6.0.12/lib' path-add 'C:/Portable/TensorRT-8.6.0.12/lib'
if $nu.os-info.family == 'windows' {
path-add 'G:\BIN\TensorRT-10.7.0.23\lib'
path-add 'G:\BIN\cudnn-windows-x86_64-9.6.0.74_cuda12-archive\bin'
}
path-add ($env.CUDA_ROOT | path join bin) path-add ($env.CUDA_ROOT | path join bin)
overlay use ../../.venv/Scripts/activate.nu overlay use ../../.venv/Scripts/activate.nu
} }
+64 -28
View File
@@ -43,10 +43,28 @@ pip_map = {
"tb-nightly": "tensorboard", "tb-nightly": "tensorboard",
"protobuf": "google.protobuf", "protobuf": "google.protobuf",
"qrcode[pil]": "qrcode", "qrcode[pil]": "qrcode",
"requirements-parser": "requirements" "requirements-parser": "requirements",
# Add more mappings as needed # Add more mappings as needed
} }
def get_node_dependencies():
restore_deps = ["basicsr"]
onnx_deps = ["onnxruntime"]
swap_deps = ["insightface"] + onnx_deps
quant_deps = ["bitsandbytes"]
io_deps = ["av"]
return {
"QrCode": ["qrcode"],
"DeepBump": onnx_deps,
"FaceSwap": swap_deps,
"LoadFaceSwapModel": swap_deps,
"LoadFaceAnalysisModel": restore_deps,
"Quantize": quant_deps,
"SaveGif": io_deps,
}
# endregion # endregion
# region ansi # region ansi
@@ -124,12 +142,12 @@ def print_formatted(text, *formats, color=None, background=None, **kwargs):
header = "[mtb install] " header = "[mtb install] "
# Handle console encoding for Unicode characters (utf-8) # Handle console encoding for Unicode characters (utf-8)
encoded_header = header.encode(sys.stdout.encoding, errors="replace").decode( encoded_header = header.encode(
sys.stdout.encoding sys.stdout.encoding, errors="replace"
) ).decode(sys.stdout.encoding)
encoded_text = formatted_text.encode(sys.stdout.encoding, errors="replace").decode( encoded_text = formatted_text.encode(
sys.stdout.encoding sys.stdout.encoding, errors="replace"
) ).decode(sys.stdout.encoding)
print( print(
" " * len(encoded_header) " " * len(encoded_header)
@@ -163,7 +181,9 @@ def run_command(cmd, ignored_lines_start=None):
try: try:
_run_command(shell_cmd, ignored_lines_start) _run_command(shell_cmd, ignored_lines_start)
except subprocess.CalledProcessError as e: except subprocess.CalledProcessError as e:
print(f"Command failed with return code: {e.returncode}", file=sys.stderr) print(
f"Command failed with return code: {e.returncode}", file=sys.stderr
)
print(e.stderr.strip(), file=sys.stderr) print(e.stderr.strip(), file=sys.stderr)
except KeyboardInterrupt: except KeyboardInterrupt:
@@ -238,7 +258,7 @@ def suppress_std():
def get_local_version(): def get_local_version():
init_file = os.path.join(os.path.dirname(__file__), "__init__.py") init_file = os.path.join(os.path.dirname(__file__), "__init__.py")
if os.path.isfile(init_file): if os.path.isfile(init_file):
with open(init_file, "r") as f: with open(init_file) as f:
tree = ast.parse(f.read()) tree = ast.parse(f.read())
for node in ast.walk(tree): for node in ast.walk(tree):
if isinstance(node, ast.Assign): if isinstance(node, ast.Assign):
@@ -256,13 +276,16 @@ def download_file(url, file_name):
with requests.get(url, stream=True) as response: with requests.get(url, stream=True) as response:
response.raise_for_status() response.raise_for_status()
total_size = int(response.headers.get("content-length", 0)) total_size = int(response.headers.get("content-length", 0))
with open(file_name, "wb") as file, tqdm( with (
desc=file_name.stem, open(file_name, "wb") as file,
total=total_size, tqdm(
unit="B", desc=file_name.stem,
unit_scale=True, total=total_size,
unit_divisor=1024, unit="B",
) as progress_bar: unit_scale=True,
unit_divisor=1024,
) as progress_bar,
):
for chunk in response.iter_content(chunk_size=8192): for chunk in response.iter_content(chunk_size=8192):
file.write(chunk) file.write(chunk)
progress_bar.update(len(chunk)) progress_bar.update(len(chunk))
@@ -302,7 +325,9 @@ def import_or_install(requirement, dry=False):
pip_install_name = pip_name + pip_spec pip_install_name = pip_name + pip_spec
if not installed: if not installed:
print_formatted(f"Installing package {pip_name}...", "italic", color="yellow") print_formatted(
f"Installing package {pip_name}...", "italic", color="yellow"
)
if dry: if dry:
print_formatted( print_formatted(
f"Dry-run: Package {pip_install_name} would be installed (import name: '{import_name}').", f"Dry-run: Package {pip_install_name} would be installed (import name: '{import_name}').",
@@ -310,7 +335,9 @@ def import_or_install(requirement, dry=False):
) )
else: else:
try: try:
run_command([executable, "-m", "pip", "install", pip_install_name]) run_command(
[executable, "-m", "pip", "install", pip_install_name]
)
print_formatted( print_formatted(
f"Package {pip_install_name} installed successfully using pip package name (import name: '{import_name}')", f"Package {pip_install_name} installed successfully using pip package name (import name: '{import_name}')",
"bold", "bold",
@@ -326,13 +353,9 @@ def import_or_install(requirement, dry=False):
def get_github_assets(tag=None): def get_github_assets(tag=None):
if tag: if tag:
tag_url = ( tag_url = f"https://api.github.com/repos/{repo_owner}/{repo_name}/releases/tags/{tag}"
f"https://api.github.com/repos/{repo_owner}/{repo_name}/releases/tags/{tag}"
)
else: else:
tag_url = ( tag_url = f"https://api.github.com/repos/{repo_owner}/{repo_name}/releases/latest"
f"https://api.github.com/repos/{repo_owner}/{repo_name}/releases/latest"
)
response = requests.get(tag_url) response = requests.get(tag_url)
if response.status_code == 404: if response.status_code == 404:
# print_formatted( # print_formatted(
@@ -361,7 +384,9 @@ except ImportError:
def main(): def main():
if len(sys.argv) == 1: if len(sys.argv) == 1:
print_formatted( print_formatted(
"mtb doesn't need an install script anymore.", "italic", color="yellow" "mtb doesn't need an install script anymore.",
"italic",
color="yellow",
) )
return return
if all(arg not in ("-p", "--path") for arg in sys.argv): if all(arg not in ("-p", "--path") for arg in sys.argv):
@@ -397,8 +422,12 @@ def main():
else: else:
repo_dir = clone_dir / repo_name repo_dir = clone_dir / repo_name
if not repo_dir.exists(): if not repo_dir.exists():
print_formatted(f"Cloning to {repo_dir}...", "italic", color="yellow") print_formatted(
run_command(["git", "clone", "--recursive", repo_url, repo_dir]) f"Cloning to {repo_dir}...", "italic", color="yellow"
)
run_command(
["git", "clone", "--recursive", repo_url, repo_dir]
)
else: else:
print_formatted( print_formatted(
f"Directory {repo_dir} already exists, we will update it..." f"Directory {repo_dir} already exists, we will update it..."
@@ -409,7 +438,14 @@ def main():
print_formatted("Checking environment...", "italic", color="yellow") print_formatted("Checking environment...", "italic", color="yellow")
missing_deps = [] missing_deps = []
install_cmd = [executable, "-m", "pip", "install", "-r", "requirements.txt"] install_cmd = [
executable,
"-m",
"pip",
"install",
"-r",
"requirements.txt",
]
run_command(install_cmd) run_command(install_cmd)
print_formatted( print_formatted(
+226 -10
View File
@@ -1,4 +1,5 @@
from io import BytesIO from io import BytesIO
from typing import Literal
import cv2 import cv2
import numpy as np import numpy as np
@@ -410,7 +411,14 @@ class MTB_BatchFloat:
RETURN_TYPES = ("FLOATS",) RETURN_TYPES = ("FLOATS",)
CATEGORY = "mtb/batch" CATEGORY = "mtb/batch"
def set_floats(self, mode, count, min, max, easing): def set_floats(
self,
mode: Literal["Steps"] | Literal["Single"] = "Steps",
count: int = 1,
min: float = 0.0, # noqa: A002
max: float = 1.0, # noqa: A002
easing: str = "Linear",
):
if mode == "Steps" and count == 1: if mode == "Steps" and count == 1:
raise ValueError( raise ValueError(
"Steps mode requires at least a count of 2 values" "Steps mode requires at least a count of 2 values"
@@ -429,6 +437,210 @@ class MTB_BatchFloat:
return (keyframes,) return (keyframes,)
class MTB_BatchSequencePlus:
"""Sequences multiple image batches with transition effects."""
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"transition": (
[
"none",
"crossfade",
"slide_left",
"slide_right",
"slide_up",
"slide_down",
"wipe_left",
"wipe_right",
"wipe_up",
"wipe_down",
"band_wipe_h",
"band_wipe_v",
],
{"default": "none"},
),
"overlap_frames": (
"INT",
{"default": 0, "min": 0, "max": 120, "step": 1},
),
"reverse": ("BOOLEAN", {"default": False}),
}
}
RETURN_TYPES = ("IMAGE",)
FUNCTION = "sequence_batches"
CATEGORY = "mtb/batch"
def apply_transition(
self,
frame1: torch.Tensor,
frame2: torch.Tensor,
transition: str,
progress: float,
):
"""Apply transition effect between two frames."""
if transition == "none":
return frame1 if progress < 0.5 else frame2
elif transition == "crossfade":
return frame1 * (1 - progress) + frame2 * progress
elif transition.startswith("slide_"):
h, w = frame1.shape[1:3]
if transition == "slide_left":
offset = int(w * progress)
frame2 = torch.roll(frame2, shifts=-offset, dims=2)
elif transition == "slide_right":
offset = int(w * progress)
frame2 = torch.roll(frame2, shifts=offset, dims=2)
elif transition == "slide_up":
offset = int(h * progress)
frame2 = torch.roll(frame2, shifts=-offset, dims=1)
elif transition == "slide_down":
offset = int(h * progress)
frame2 = torch.roll(frame2, shifts=offset, dims=1)
return frame1 * (1 - progress) + frame2 * progress
elif transition.startswith("wipe_"):
h, w = frame1.shape[1:3]
mask = torch.zeros_like(frame1)
if transition == "wipe_left":
edge = int(w * progress)
mask[:, :, :edge, :] = 1
elif transition == "wipe_right":
edge = int(w * (1 - progress))
mask[:, :, edge:, :] = 1
elif transition == "wipe_up":
edge = int(h * progress)
mask[:, :edge, :, :] = 1
elif transition == "wipe_down":
edge = int(h * (1 - progress))
mask[:, edge:, :, :] = 1
return frame1 * (1 - mask) + frame2 * mask
elif transition.startswith("band_wipe_"):
h, w = frame1.shape[1:3]
mask = torch.zeros_like(frame1)
num_bands = 10 # Number of bands
if transition == "band_wipe_h":
band_width = w / num_bands
for i in range(num_bands):
edge = int((w * progress) - (i * band_width))
start = int(i * band_width)
end = int(min(start + edge, (i + 1) * band_width))
if end > start:
mask[:, :, start:end, :] = 1
else: # band_wipe_v
band_height = h / num_bands
for i in range(num_bands):
edge = int((h * progress) - (i * band_height))
start = int(i * band_height)
end = int(min(start + edge, (i + 1) * band_height))
if end > start:
mask[:, start:end, :, :] = 1
return frame1 * (1 - mask) + frame2 * mask
return frame1
def sequence_batches(
self, transition: str, overlap_frames: int, reverse: bool, **kwargs
):
images: list[torch.Tensor] = list(kwargs.values())
if reverse:
images = images[::-1]
processed_images: list[torch.Tensor] = []
for img in images:
if len(img.shape) == 3:
img = img.unsqueeze(0)
processed_images.append(img)
if overlap_frames == 0 or transition == "none":
return (torch.cat(processed_images, dim=0),)
result_frames: list[torch.Tensor] = []
if len(processed_images) > 0:
result_frames.extend(
list(processed_images[0][: -overlap_frames // 2])
)
for i in range(1, len(processed_images)):
prev_batch = processed_images[i - 1]
curr_batch = processed_images[i]
prev_frames = min(overlap_frames // 2, len(prev_batch))
next_frames = min(overlap_frames // 2, len(curr_batch))
total_overlap = prev_frames + next_frames
if total_overlap < 2:
# when not enough frames for transition, just concatenate
result_frames.extend(list(prev_batch[-prev_frames:]))
result_frames.extend(list(curr_batch[:next_frames]))
continue
for t in range(total_overlap):
progress = t / (total_overlap - 1)
prev_idx = (
len(prev_batch) - prev_frames + min(t, prev_frames - 1)
)
next_idx = max(0, t - prev_frames)
transition_frame = self.apply_transition(
prev_batch[prev_idx : prev_idx + 1],
curr_batch[next_idx : next_idx + 1],
transition,
progress,
)
result_frames.append(transition_frame[0])
if i < len(processed_images) - 1:
result_frames.extend(
list(curr_batch[next_frames : -overlap_frames // 2])
)
else:
result_frames.extend(list(curr_batch[next_frames:]))
result = torch.stack(result_frames, dim=0)
return (result,)
class MTB_BatchSequence:
"""Sequences multiple image batches one after another"""
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"reverse": ("BOOLEAN", {"default": False}),
}
}
RETURN_TYPES = ("IMAGE",)
FUNCTION = "sequence_batches"
CATEGORY = "mtb/batch"
def sequence_batches(self, reverse: bool, **kwargs):
images = list(kwargs.values())
if reverse:
images = images[::-1]
processed = []
for img in images:
if len(img.shape) == 3:
img = img.unsqueeze(0)
processed.append(img)
return (torch.cat(processed, dim=0),)
class MTB_BatchMerge: class MTB_BatchMerge:
"""Merges multiple image batches with different frame counts""" """Merges multiple image batches with different frame counts"""
@@ -711,7 +923,9 @@ class MTB_PlotBatchFloat:
ax.set_xlim(1, max_length) # Set X-axis limits ax.set_xlim(1, max_length) # Set X-axis limits
np.random.seed(seed) np.random.seed(seed)
colors = np.random.rand(len(kwargs), 3) # Generate random RGB values colors = np.random.rand(len(kwargs), 3) # Generate random RGB values
for color, (label, values) in zip(colors, kwargs.items()): for color, (label, values) in zip(
colors, kwargs.items(), strict=False
):
ax.plot(x_values[: len(values)], values, label=label, color=color) ax.plot(x_values[: len(values)], values, label=label, color=color)
ax.legend( ax.legend(
title="Legend", title="Legend",
@@ -1026,17 +1240,19 @@ class MTB_BatchShake:
__nodes__ = [ __nodes__ = [
MTB_BatchFloat,
MTB_Batch2dTransform, MTB_Batch2dTransform,
MTB_BatchShape, MTB_BatchFloat,
MTB_BatchMake,
MTB_BatchFloatAssemble, MTB_BatchFloatAssemble,
MTB_BatchFloatFill, MTB_BatchFloatFill,
MTB_BatchFloatNormalize,
MTB_BatchMerge,
MTB_BatchShake,
MTB_PlotBatchFloat,
MTB_BatchTimeWrap,
MTB_BatchFloatFit, MTB_BatchFloatFit,
MTB_BatchFloatMath, MTB_BatchFloatMath,
MTB_BatchFloatNormalize,
MTB_BatchMake,
MTB_BatchMerge,
MTB_BatchSequence,
MTB_BatchSequencePlus,
MTB_BatchShake,
MTB_BatchShape,
MTB_BatchTimeWrap,
MTB_PlotBatchFloat,
] ]
+123 -1
View File
@@ -3,10 +3,127 @@ import shutil
from pathlib import Path from pathlib import Path
import folder_paths import folder_paths
import torch
from ..log import log from ..log import log
from ..utils import here from ..utils import here
Conditioning = list[tuple[torch.Tensor, dict[str, torch.Tensor]]]
def check_condition(conditioning: Conditioning):
has_cn = False
if len(conditioning) > 1:
log.warn(
"More than one conditioning was provided. Only the first one will be used."
)
first = conditioning[0]
cond, kwargs = first
log.debug("Conditioning Shape")
log.debug(cond.shape)
log.debug("Conditioning keys")
log.debug([f"\t{k} - {type(kwargs[k])}" for k in kwargs])
if "control" in kwargs:
log.debug("Conditioning contains a controlnet")
has_cn = True
if "pooled_output" not in kwargs:
raise ValueError(
"Conditioning is not valid. Missing 'pooled_output' key."
)
return has_cn
class MTB_InterpolateCondition:
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"blend": (
"FLOAT",
{"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.01},
),
},
}
RETURN_TYPES = ("CONDITIONING",)
CATEGORY = "mtb/conditioning"
FUNCTION = "execute"
def execute(
self, blend: float, **kwargs: Conditioning
) -> tuple[Conditioning]:
blend = max(0.0, min(1.0, blend))
conditions: list[Conditioning] = list(kwargs.values())
num_conditions = len(conditions)
if num_conditions < 2:
raise ValueError("At least two conditioning inputs are required.")
segment_length = 1.0 / (num_conditions - 1)
segment_index = min(int(blend // segment_length), num_conditions - 2)
local_blend = (
blend - (segment_index * segment_length)
) / segment_length
cond_from = conditions[segment_index]
cond_to = conditions[segment_index + 1]
from_cn = check_condition(cond_from)
to_cn = check_condition(cond_to)
if from_cn and to_cn:
raise ValueError(
"Interpolating conditions cannot both contain ControlNets"
)
try:
interpolated_condition = [
(1.0 - local_blend) * c_from + local_blend * c_to
for c_from, c_to in zip(
cond_from[0][0], cond_to[0][0], strict=False
)
]
except Exception as e:
print(f"Error during interpolation: {e}")
raise
pooled_from = cond_from[0][1].get(
"pooled_output",
torch.zeros_like(
next(iter(cond_from[0][1].values()), torch.tensor([]))
),
)
pooled_to = cond_to[0][1].get(
"pooled_output",
torch.zeros_like(
next(iter(cond_from[0][1].values()), torch.tensor([]))
),
)
interpolated_pooled = (
1.0 - local_blend
) * pooled_from + local_blend * pooled_to
res = {"pooled_output": interpolated_pooled}
if from_cn:
res["control"] = cond_from[0][1]["control"]
res["control_apply_to_uncond"] = cond_from[0][1][
"control_apply_to_uncond"
]
if to_cn:
res["control"] = cond_to[0][1]["control"]
res["control_apply_to_uncond"] = cond_to[0][1][
"control_apply_to_uncond"
]
return ([(torch.stack(interpolated_condition), res)],)
class MTB_InterpolateClipSequential: class MTB_InterpolateClipSequential:
@classmethod @classmethod
@@ -213,4 +330,9 @@ class MTB_StylesLoader:
return (self.options[style_name][0], self.options[style_name][1]) return (self.options[style_name][0], self.options[style_name][1])
__nodes__ = [MTB_SmartStep, MTB_StylesLoader, MTB_InterpolateClipSequential] __nodes__ = [
MTB_SmartStep,
MTB_StylesLoader,
MTB_InterpolateClipSequential,
MTB_InterpolateCondition,
]
+38 -1
View File
@@ -59,6 +59,36 @@ class MTB_SplitBbox:
return (bbox[0], bbox[1], bbox[2], bbox[3]) return (bbox[0], bbox[1], bbox[2], bbox[3])
class MTB_UpscaleBboxBy:
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"bbox": ("BBOX",),
"scale": ("FLOAT", {"default": 1.0}),
},
}
CATEGORY = "mtb/crop"
RETURN_TYPES = ("BBOX",)
FUNCTION = "upscale"
def upscale(
self, bbox: tuple[int, int, int, int], scale: float
) -> tuple[tuple[int, int, int, int]]:
x, y, width, height = bbox
# scaled = (x * scale, y * scale, width * scale, height * scale)
scaled = (
int(x * scale),
int(y * scale),
int(width * scale),
int(height * scale),
)
return (scaled,)
class MTB_BboxFromMask: class MTB_BboxFromMask:
"""From a mask extract the bounding box""" """From a mask extract the bounding box"""
@@ -342,4 +372,11 @@ class MTB_Uncrop:
return (pil2tensor(out_images),) return (pil2tensor(out_images),)
__nodes__ = [MTB_BboxFromMask, MTB_Bbox, MTB_Crop, MTB_Uncrop, MTB_SplitBbox] __nodes__ = [
MTB_BboxFromMask,
MTB_Bbox,
MTB_Crop,
MTB_Uncrop,
MTB_SplitBbox,
MTB_UpscaleBboxBy,
]
+2
View File
@@ -78,6 +78,7 @@ class MTB_LoadFaceEnhanceModel:
RETURN_NAMES = ("model",) RETURN_NAMES = ("model",)
FUNCTION = "load_model" FUNCTION = "load_model"
CATEGORY = "mtb/facetools" CATEGORY = "mtb/facetools"
DEPRECATED = True
def load_model(self, model_name, upscale=2, bg_upsampler=None): def load_model(self, model_name, upscale=2, bg_upsampler=None):
from gfpgan import GFPGANer from gfpgan import GFPGANer
@@ -163,6 +164,7 @@ class MTB_RestoreFace:
RETURN_TYPES = ("IMAGE",) RETURN_TYPES = ("IMAGE",)
FUNCTION = "restore" FUNCTION = "restore"
CATEGORY = "mtb/facetools" CATEGORY = "mtb/facetools"
DEPRECATED = True
@classmethod @classmethod
def INPUT_TYPES(cls): def INPUT_TYPES(cls):
+3
View File
@@ -40,6 +40,7 @@ class MTB_LoadFaceAnalysisModel:
RETURN_TYPES = ("FACE_ANALYSIS_MODEL",) RETURN_TYPES = ("FACE_ANALYSIS_MODEL",)
FUNCTION = "load_model" FUNCTION = "load_model"
CATEGORY = "mtb/facetools" CATEGORY = "mtb/facetools"
DEPRECATED = True
def load_model(self, faceswap_model: str): def load_model(self, faceswap_model: str):
if faceswap_model == "antelopev2": if faceswap_model == "antelopev2":
@@ -77,6 +78,7 @@ class MTB_LoadFaceSwapModel:
RETURN_TYPES = ("FACESWAP_MODEL",) RETURN_TYPES = ("FACESWAP_MODEL",)
FUNCTION = "load_model" FUNCTION = "load_model"
CATEGORY = "mtb/facetools" CATEGORY = "mtb/facetools"
DEPRECATED = True
def load_model(self, faceswap_model: str): def load_model(self, faceswap_model: str):
model_path = get_model_path("insightface", faceswap_model) model_path = get_model_path("insightface", faceswap_model)
@@ -126,6 +128,7 @@ class MTB_FaceSwap:
RETURN_TYPES = ("IMAGE",) RETURN_TYPES = ("IMAGE",)
FUNCTION = "swap" FUNCTION = "swap"
CATEGORY = "mtb/facetools" CATEGORY = "mtb/facetools"
DEPRECATED = True
def swap( def swap(
self, self,
+11 -4
View File
@@ -1,5 +1,4 @@
from pathlib import Path from pathlib import Path
from typing import List
import comfy import comfy
import comfy.model_management as model_management import comfy.model_management as model_management
@@ -15,10 +14,13 @@ from ..utils import get_model_path
class MTB_LoadFilmModel: class MTB_LoadFilmModel:
"""Loads a FILM model""" """Loads a FILM model
[DEPRECATED] Use ComfyUI-FrameInterpolation instead
"""
@staticmethod @staticmethod
def get_models() -> List[Path]: def get_models() -> list[Path]:
models_paths = get_model_path("FILM").iterdir() models_paths = get_model_path("FILM").iterdir()
return [x for x in models_paths if x.suffix in [".onnx", ".pth"]] return [x for x in models_paths if x.suffix in [".onnx", ".pth"]]
@@ -37,6 +39,7 @@ class MTB_LoadFilmModel:
RETURN_TYPES = ("FILM_MODEL",) RETURN_TYPES = ("FILM_MODEL",)
FUNCTION = "load_model" FUNCTION = "load_model"
CATEGORY = "mtb/frame iterpolation" CATEGORY = "mtb/frame iterpolation"
DEPRECATED = True
def load_model(self, film_model: str): def load_model(self, film_model: str):
model_path = get_model_path("FILM", film_model) model_path = get_model_path("FILM", film_model)
@@ -56,7 +59,10 @@ class MTB_LoadFilmModel:
class MTB_FilmInterpolation: class MTB_FilmInterpolation:
"""Google Research FILM frame interpolation for large motion""" """Google Research FILM frame interpolation for large motion
[DEPRECATED] Use ComfyUI-FrameInterpolation instead
"""
@classmethod @classmethod
def INPUT_TYPES(cls): def INPUT_TYPES(cls):
@@ -71,6 +77,7 @@ class MTB_FilmInterpolation:
RETURN_TYPES = ("IMAGE",) RETURN_TYPES = ("IMAGE",)
FUNCTION = "do_interpolation" FUNCTION = "do_interpolation"
CATEGORY = "mtb/frame iterpolation" CATEGORY = "mtb/frame iterpolation"
DEPRECATED = True
def do_interpolation( def do_interpolation(
self, self,
+1 -3
View File
@@ -702,7 +702,6 @@ class MTB_Blur:
) )
blurred_images.append(blurred) blurred_images.append(blurred)
image_np = np.array(blurred_images)
else: else:
for i in range(image.size(0)): for i in range(image.size(0)):
blurred = gaussian( blurred = gaussian(
@@ -710,8 +709,7 @@ class MTB_Blur:
) )
blurred_images.append(blurred) blurred_images.append(blurred)
image_np = np.array(blurred_images) return (np2tensor(blurred_images),)
return (np2tensor(image_np).squeeze(0),)
class MTB_Sharpen: class MTB_Sharpen:
+168 -24
View File
@@ -1,4 +1,11 @@
import json
import os
import numpy as np
import torch import torch
from comfy.cli_args import args
from PIL import Image
from PIL.PngImagePlugin import PngInfo
from ..log import log from ..log import log
@@ -8,13 +15,21 @@ class MTB_StackImages:
@classmethod @classmethod
def INPUT_TYPES(cls): def INPUT_TYPES(cls):
return {"required": {"vertical": ("BOOLEAN", {"default": False})}} return {
"required": {"vertical": ("BOOLEAN", {"default": False})},
"optional": {
"match_method": (
["error", "smallest", "largest"],
{"default": "error"},
)
},
}
RETURN_TYPES = ("IMAGE",) RETURN_TYPES = ("IMAGE",)
FUNCTION = "stack" FUNCTION = "stack"
CATEGORY = "mtb/image utils" CATEGORY = "mtb/image utils"
def stack(self, vertical, **kwargs): def stack(self, vertical, match_method="error", **kwargs):
if not kwargs: if not kwargs:
raise ValueError("At least one tensor must be provided.") raise ValueError("At least one tensor must be provided.")
@@ -32,23 +47,50 @@ class MTB_StackImages:
self.duplicate_frames(tensor, max_batch_size) self.duplicate_frames(tensor, max_batch_size)
for tensor in normalized_tensors for tensor in normalized_tensors
] ]
if match_method != "error":
if vertical: if vertical:
width = normalized_tensors[0].shape[2] # match widths
if any(tensor.shape[2] != width for tensor in normalized_tensors): widths = [tensor.shape[2] for tensor in normalized_tensors]
raise ValueError( target_width = (
"All tensors must have the same width " min(widths) if match_method == "smallest" else max(widths)
"for vertical stacking."
) )
dim = 1 normalized_tensors = [
self.resize_tensor(tensor, width=target_width)
for tensor in normalized_tensors
]
else:
# match heights
heights = [tensor.shape[1] for tensor in normalized_tensors]
target_height = (
min(heights)
if match_method == "smallest"
else max(heights)
)
normalized_tensors = [
self.resize_tensor(tensor, height=target_height)
for tensor in normalized_tensors
]
else: else:
height = normalized_tensors[0].shape[1] if vertical:
if any(tensor.shape[1] != height for tensor in normalized_tensors): width = normalized_tensors[0].shape[2]
raise ValueError( if any(
"All tensors must have the same height " tensor.shape[2] != width for tensor in normalized_tensors
"for horizontal stacking." ):
) raise ValueError(
dim = 2 "All tensors must have the same width "
"for vertical stacking."
)
else:
height = normalized_tensors[0].shape[1]
if any(
tensor.shape[1] != height for tensor in normalized_tensors
):
raise ValueError(
"All tensors must have the same height "
"for horizontal stacking."
)
dim = 1 if vertical else 2
stacked_tensor = torch.cat(normalized_tensors, dim=dim) stacked_tensor = torch.cat(normalized_tensors, dim=dim)
@@ -64,7 +106,7 @@ class MTB_StackImages:
elif channels == 3: elif channels == 3:
alpha_channel = torch.ones( alpha_channel = torch.ones(
tensor.shape[:-1] + (1,), device=tensor.device tensor.shape[:-1] + (1,), device=tensor.device
) # Add an alpha channel )
return torch.cat((tensor, alpha_channel), dim=-1) return torch.cat((tensor, alpha_channel), dim=-1)
else: else:
raise ValueError( raise ValueError(
@@ -87,6 +129,30 @@ class MTB_StackImages:
else: else:
return tensor return tensor
def resize_tensor(self, tensor, width=None, height=None):
"""Resize tensor to specified width or height while maintaining aspect ratio."""
current_height, current_width = tensor.shape[1:3]
if width is not None and width != current_width:
scale_factor = width / current_width
new_height = int(current_height * scale_factor)
new_width = width
elif height is not None and height != current_height:
scale_factor = height / current_height
new_width = int(current_width * scale_factor)
new_height = height
else:
return tensor
resized = torch.nn.functional.interpolate(
tensor.permute(0, 3, 1, 2),
size=(new_height, new_width),
mode="bilinear",
align_corners=False,
)
return resized.permute(0, 2, 3, 1)
class MTB_PickFromBatch: class MTB_PickFromBatch:
"""Pick a specific number of images from a batch. """Pick a specific number of images from a batch.
@@ -113,11 +179,6 @@ class MTB_PickFromBatch:
# Limit count to the available number of images in the batch # Limit count to the available number of images in the batch
count = min(count, batch_size) count = min(count, batch_size)
if count < batch_size:
log.warning(
f"Requested {count} images, "
f"but only {batch_size} are available."
)
if from_direction == "end": if from_direction == "end":
selected_tensors = image[-count:] selected_tensors = image[-count:]
@@ -127,4 +188,87 @@ class MTB_PickFromBatch:
return (selected_tensors,) return (selected_tensors,)
__nodes__ = [MTB_StackImages, MTB_PickFromBatch] import folder_paths
class MTB_SaveImage:
def __init__(self):
self.output_dir = folder_paths.get_output_directory()
self.type = "output"
self.prefix_append = ""
self.compress_level = 4
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"images": ("IMAGE", {"tooltip": "The images to save."}),
"filename_prefix": (
"STRING",
{
"default": "ComfyUI",
"tooltip": "The prefix for the file to save. This may include formatting information such as %date:yyyy-MM-dd% or %Empty Latent Image.width% to include values from nodes.",
},
),
},
"hidden": {"prompt": "PROMPT", "extra_pnginfo": "EXTRA_PNGINFO"},
}
RETURN_TYPES = ("IMAGE",)
FUNCTION = "save_images"
# OUTPUT_NODE = True
CATEGORY = "mtb/image utils"
DESCRIPTION = """Saves the input images to your ComfyUI output directory.
This behaves exactly like the native SaveImage node but isn't an output node.
The reason I made this is to allow 'inlining' image save in loops for instance,
using the native node there wouldn't run for each iteration of the loop."""
def save_images(
self,
images,
filename_prefix="ComfyUI",
prompt=None,
extra_pnginfo=None,
):
filename_prefix += self.prefix_append
full_output_folder, filename, counter, subfolder, filename_prefix = (
folder_paths.get_save_image_path(
filename_prefix,
self.output_dir,
images[0].shape[1],
images[0].shape[0],
)
)
results = list()
for batch_number, image in enumerate(images):
i = 255.0 * image.cpu().numpy()
img = Image.fromarray(np.clip(i, 0, 255).astype(np.uint8))
metadata = None
if not args.disable_metadata:
metadata = PngInfo()
if prompt is not None:
metadata.add_text("prompt", json.dumps(prompt))
if extra_pnginfo is not None:
for x in extra_pnginfo:
metadata.add_text(x, json.dumps(extra_pnginfo[x]))
filename_with_batch_num = filename.replace(
"%batch_num%", str(batch_number)
)
file = f"{filename_with_batch_num}_{counter:05}_.png"
img.save(
os.path.join(full_output_folder, file),
pnginfo=metadata,
compress_level=self.compress_level,
)
results.append(
{"filename": file, "subfolder": subfolder, "type": self.type}
)
counter += 1
return {"ui": {"images": results}, "result": (images,)}
__nodes__ = [MTB_StackImages, MTB_PickFromBatch, MTB_SaveImage]
+55 -19
View File
@@ -2,9 +2,9 @@ import json
import subprocess import subprocess
import uuid import uuid
from pathlib import Path from pathlib import Path
from typing import List, Optional
import comfy.model_management as model_management import comfy.model_management as model_management
import comfy.utils
import folder_paths import folder_paths
import numpy as np import numpy as np
import torch import torch
@@ -41,6 +41,7 @@ class MTB_ReadPlaylist:
RETURN_TYPES = ("PLAYLIST",) RETURN_TYPES = ("PLAYLIST",)
FUNCTION = "read_playlist" FUNCTION = "read_playlist"
CATEGORY = "mtb/IO" CATEGORY = "mtb/IO"
EXPERIMENTAL = True
def read_playlist( def read_playlist(
self, self,
@@ -83,6 +84,7 @@ class MTB_AddToPlaylist:
OUTPUT_NODE = True OUTPUT_NODE = True
FUNCTION = "add_to_playlist" FUNCTION = "add_to_playlist"
CATEGORY = "mtb/IO" CATEGORY = "mtb/IO"
EXPERIMENTAL = True
def add_to_playlist( def add_to_playlist(
self, self,
@@ -117,7 +119,10 @@ class MTB_AddToPlaylist:
class MTB_ExportWithFfmpeg: class MTB_ExportWithFfmpeg:
"""Export with FFmpeg (Experimental)""" """Export with FFmpeg (Experimental).
[DEPRACATED] Use VHS nodes instead
"""
@classmethod @classmethod
def INPUT_TYPES(cls): def INPUT_TYPES(cls):
@@ -143,6 +148,7 @@ class MTB_ExportWithFfmpeg:
RETURN_TYPES = ("VIDEO",) RETURN_TYPES = ("VIDEO",)
OUTPUT_NODE = True OUTPUT_NODE = True
FUNCTION = "export_prores" FUNCTION = "export_prores"
DEPRECATED = True
CATEGORY = "mtb/IO" CATEGORY = "mtb/IO"
def export_prores( def export_prores(
@@ -151,10 +157,9 @@ class MTB_ExportWithFfmpeg:
prefix: str, prefix: str,
format: str, format: str,
codec: str, codec: str,
images: Optional[torch.Tensor] = None, images: torch.Tensor | None = None,
playlist: Optional[List[str]] = None, playlist: list[str] | None = None,
): ):
pix_fmt = "rgb48le" if codec == "prores_ks" else "yuv420p"
file_ext = format file_ext = format
file_id = f"{prefix}_{uuid.uuid4()}.{file_ext}" file_id = f"{prefix}_{uuid.uuid4()}.{file_ext}"
@@ -208,9 +213,11 @@ class MTB_ExportWithFfmpeg:
frames = tensor2np(images) frames = tensor2np(images)
log.debug(f"Frames type {type(frames[0])}") log.debug(f"Frames type {type(frames[0])}")
log.debug(f"Exporting {len(frames)} frames") log.debug(f"Exporting {len(frames)} frames")
height, width, channels = frames[0].shape
has_alpha = channels == 4
out_path = (output_dir / file_id).as_posix()
if codec == "gif": if codec == "gif":
out_path = (output_dir / file_id).as_posix()
command = [ command = [
"ffmpeg", "ffmpeg",
"-f", "-f",
@@ -233,12 +240,28 @@ class MTB_ExportWithFfmpeg:
process.stdin.close() process.stdin.close()
process.wait() process.wait()
return (out_path,)
else: else:
frames = [frame.astype(np.uint16) * 257 for frame in frames] if has_alpha:
if codec in ["prores_ks", "libx264", "libx265"]:
height, width, _ = frames[0].shape pix_fmt = (
"yuva444p" if codec == "prores_ks" else "yuva420p"
out_path = (output_dir / file_id).as_posix() )
frames = [
frame.astype(np.uint16) * 257 for frame in frames
]
else:
log.warning(
f"Alpha channel not supported for codec {codec}. Alpha will be ignored."
)
frames = [
frame[:, :, :3].astype(np.uint16) * 257
for frame in frames
]
pix_fmt = "rgb48le" if codec == "prores_ks" else "yuv420p"
else:
pix_fmt = "rgb48le" if codec == "prores_ks" else "yuv420p"
frames = [frame.astype(np.uint16) * 257 for frame in frames]
# Prepare the FFmpeg command # Prepare the FFmpeg command
command = [ command = [
@@ -258,17 +281,26 @@ class MTB_ExportWithFfmpeg:
"-", "-",
"-c:v", "-c:v",
codec, codec,
"-r",
str(fps),
"-y",
out_path,
] ]
if codec == "prores_ks":
command.extend(["-profile:v", "4444"])
command.extend(
[
"-r",
str(fps),
"-y",
out_path,
]
)
process = subprocess.Popen(command, stdin=subprocess.PIPE) process = subprocess.Popen(command, stdin=subprocess.PIPE)
pbar = comfy.utils.ProgressBar(len(frames))
for frame in frames: for frame in frames:
model_management.throw_exception_if_processing_interrupted()
process.stdin.write(frame.tobytes()) process.stdin.write(frame.tobytes())
pbar.update(1)
process.stdin.close() process.stdin.close()
process.wait() process.wait()
@@ -280,9 +312,9 @@ def prepare_animated_batch(
batch: torch.Tensor, batch: torch.Tensor,
pingpong=False, pingpong=False,
resize_by=1.0, resize_by=1.0,
resample_filter: Optional[Image.Resampling] = None, resample_filter: Image.Resampling | None = None,
image_type=np.uint8, image_type=np.uint8,
) -> List[Image.Image]: ) -> list[Image.Image]:
images = tensor2np(batch) images = tensor2np(batch)
images = [frame.astype(image_type) for frame in images] images = [frame.astype(image_type) for frame in images]
@@ -308,7 +340,10 @@ def prepare_animated_batch(
# todo: deprecate for apng # todo: deprecate for apng
class MTB_SaveGif: class MTB_SaveGif:
"""Save the images from the batch as a GIF""" """Save the images from the batch as a GIF.
[DEPRACATED] Use VHS nodes instead
"""
@classmethod @classmethod
def INPUT_TYPES(cls): def INPUT_TYPES(cls):
@@ -328,6 +363,7 @@ class MTB_SaveGif:
OUTPUT_NODE = True OUTPUT_NODE = True
CATEGORY = "mtb/IO" CATEGORY = "mtb/IO"
FUNCTION = "save_gif" FUNCTION = "save_gif"
DEPRECATED = True
def save_gif( def save_gif(
self, self,
+161
View File
@@ -0,0 +1,161 @@
import os
import subprocess
import tempfile
import numpy as np
import torch
from PIL import Image
from ..log import log
class ImageH264Compression:
"""Encodes the input with h264 compression using a configurable CRF."""
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"image": (
"IMAGE",
{
"tooltip": "The input image tensor to be compressed and decompressed."
},
),
"crf": (
"INT",
{
"default": 23,
"min": 0,
"max": 51,
"step": 1,
"tooltip": "Constant Rate Factor for h264 encoding (lower values mean higher quality).",
},
),
}
}
RETURN_TYPES = ("IMAGE",)
FUNCTION = "compress_and_decompress"
CATEGORY = "image"
DESCRIPTION = """
**Encodes the input with h264 compression using a configurable CRF**.
> [!IMPORTANT]
> This node is not really needed with the latest version of LTXVideo.
> [!NOTE]
> This was recommended by the creators of LTX over banodoco's discord.
*Orginal code from [mix](https://github.com/XmYx)*"""
def _compress_decompress_ffmpeg(self, img_array, crf):
with tempfile.TemporaryDirectory() as temp_dir:
input_path = os.path.join(temp_dir, "input.png")
output_path = os.path.join(temp_dir, "output.mp4")
decoded_path = os.path.join(temp_dir, "decoded.png")
Image.fromarray(img_array).save(input_path)
encode_command = [
"ffmpeg",
"-y",
"-i",
input_path,
"-c:v",
"libx264",
"-crf",
str(crf),
"-pix_fmt",
"yuv420p",
"-frames:v",
"1",
output_path,
]
subprocess.run(encode_command, capture_output=True)
decode_command = [
"ffmpeg",
"-y",
"-i",
output_path,
"-frames:v",
"1",
decoded_path,
]
subprocess.run(decode_command, capture_output=True)
decoded_img = np.array(Image.open(decoded_path))
return decoded_img
def compress_and_decompress(self, image, crf):
import io
output_images = []
try:
import av
for img_tensor in image:
img_array = img_tensor.cpu().numpy()
img_array = (img_array * 255).astype(np.uint8)
img_array = img_array.copy(
order="C"
) # Ensure contiguous array
output = io.BytesIO()
# Encode the image to h264 with the given CRF
container = av.open(output, mode="w", format="mp4")
stream = container.add_stream("h264", rate=1)
stream.width = img_array.shape[1]
stream.height = img_array.shape[0]
stream.pix_fmt = "yuv420p"
stream.options = {"crf": str(crf)}
frame = av.VideoFrame.from_ndarray(img_array, format="rgb24")
for packet in stream.encode(frame):
container.mux(packet)
for packet in stream.encode():
container.mux(packet)
container.close()
# Decode the video back to an image
output.seek(0)
container = av.open(output, mode="r", format="mp4")
decoded_frames = []
for frame in container.decode(video=0):
img_decoded = frame.to_ndarray(format="rgb24")
decoded_frames.append(img_decoded)
container.close()
if len(decoded_frames) > 0:
img_decoded = decoded_frames[0]
img_decoded = torch.from_numpy(
img_decoded.astype(np.float32) / 255.0
)
output_images.append(img_decoded)
else:
# If decoding failed, use the original image
output_images.append(img_tensor)
except ImportError:
log.warning(
"PyAv is not installed... Falling back to the ffmpeg cli"
)
for img_tensor in image:
img_array = (img_tensor.cpu().numpy() * 255).astype(np.uint8)
decoded_img = self._compress_decompress_ffmpeg(img_array, crf)
img_decoded = torch.from_numpy(
decoded_img.astype(np.float32) / 255.0
)
output_images.append(img_decoded)
output_images = torch.stack(output_images).to(image.device)
return (output_images,)
# fmt: off
__nodes__ = [
ImageH264Compression
]
+2 -1
View File
@@ -1,6 +1,5 @@
import comfy.utils import comfy.utils
from PIL import Image from PIL import Image
from rembg import remove
from ..utils import pil2tensor, tensor2pil from ..utils import pil2tensor, tensor2pil
@@ -64,6 +63,8 @@ class MTB_ImageRemoveBackgroundRembg:
post_process_mask, post_process_mask,
bgcolor, bgcolor,
): ):
from rembg import remove
pbar = comfy.utils.ProgressBar(image.size(0)) pbar = comfy.utils.ProgressBar(image.size(0))
images = tensor2pil(image) images = tensor2pil(image)
+351
View File
@@ -0,0 +1,351 @@
import os
import subprocess
import tempfile
import comfy.utils
import torch
from ..log import log
from ..utils import nextAvailable, tensor2pil
RELATIVE_NOTICE = """
Absolute paths are kept as is, relatives are from the output directory.
"""
class MTB_PostshotTrain:
CATEGORY = "mtb/postshot"
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"images": (
"IMAGE",
{"tooltip": "These image will get save to disk first"},
),
"profile": (
[
"NeRF L",
"NeRF M",
"NeRF S",
"NeRF XL",
"NeRF XXL",
"Splat ADC",
"Splat MCMC",
],
{
"default": "Splat MCMC",
"tooltip": "The radiance field model profile to train",
},
),
"image_select": (
["all", "best"],
{
"default": "best",
"tooltip": "How to select training images from the source image sets",
},
),
"train_steps_limit": (
"INT",
{
"default": 30,
"min": 1,
"max": 1000,
"tooltip": "Number of kSteps to train the model for",
},
),
"output_path": (
"STRING",
{
"default": "output",
"tooltip": (
"path to save the project to" f"{RELATIVE_NOTICE}"
),
},
),
"postshot_cli": (
"STRING",
{
"default": "C:/Program Files/Jawset Postshot/bin/postshot-cli.exe"
},
),
},
"optional": {
"gpu": (
"INT",
{
"default": 0,
"min": 0,
"max": 255,
"tooltip": "Specify the index of the GPU to use",
},
),
"num_train_images": (
"INT",
{
"default": 0,
"min": 0,
"tooltip": "If image-select best is used, specifies the number of training images to select",
},
),
"max_image_size": (
"INT",
{
"default": 1600,
"min": 0,
"tooltip": "Downscale training images such that their longer edge is at most this value in pixels. Disabled if zero.",
},
),
"max_num_features": (
"INT",
{
"default": 8,
"min": 1,
"tooltip": "Maximum number of 2D kFeatures extracted from each image.",
},
),
"splat_density": (
"FLOAT",
{
"default": 1.0,
"min": 0.125,
"max": 8.0,
"tooltip": (
"Controls how much additional splats "
"are generated during training."
"Applies only in 'Splat ADC' profile."
),
},
),
"max_num_splats": (
"INT",
{
"default": 3000,
"min": 1,
"tooltip": (
"Sets the maximum number of splats (in kSplats)"
" created during training. "
"Applies only in 'Splat MCMC' profile."
),
},
),
"export_splat_ply": (
"STRING",
{
"default": "",
"tooltip": (
"If not empty will also save a ply file."
f"{RELATIVE_NOTICE}"
),
},
),
},
}
RETURN_TYPES = ("STRING",)
OUTPUT_NODE = True
RETURN_NAMES = ("project_file_path",)
FUNCTION = "train_model"
def train_model(
self,
images: torch.Tensor,
profile: str,
image_select: str,
train_steps_limit: int,
output_path: str,
gpu=0,
num_train_images=0,
max_image_size=1600,
max_num_features=8,
splat_density=1.0,
max_num_splats=3000,
export_splat_ply="",
postshot_cli="",
):
if not output_path.endswith(".psht"):
output_path += ".psht"
output_path = nextAvailable(output_path)
output_path.parent.mkdir(exist_ok=True)
pbar = comfy.utils.ProgressBar(200 + images.size(0))
try:
with tempfile.TemporaryDirectory() as temp_dir:
image_paths = []
images_pil = tensor2pil(images)
for i, img in enumerate(images_pil):
try:
img_path = os.path.join(temp_dir, f"image_{i:04d}.png")
img.save(img_path)
image_paths.append(img_path)
except Exception as e:
raise RuntimeError(
f"Failed to save image {i}: {str(e)}"
) from e
pbar.update(1)
if not image_paths:
raise ValueError("No valid images to process")
cmd = [postshot_cli, "train"]
for img_path in image_paths:
cmd.extend(["-i", img_path])
cmd.extend(
[
"-p",
profile,
"--image-select",
image_select,
"-s",
str(train_steps_limit),
"-o",
output_path.as_posix(),
]
)
if gpu is not None:
cmd.extend(["--gpu", str(gpu)])
if num_train_images > 0 and image_select == "best":
cmd.extend(["--num-train-images", str(num_train_images)])
if max_image_size > 0:
cmd.extend(["--max-image-size", str(max_image_size)])
if max_num_features != 8:
cmd.extend(["--max-num-features", str(max_num_features)])
if profile == "Splat ADC" and splat_density != 1.0:
cmd.extend(["--splat-density", str(splat_density)])
if profile == "Splat MCMC" and max_num_splats != 3000:
cmd.extend(["--max-num-splats", str(max_num_splats)])
if export_splat_ply:
export_splat_ply = nextAvailable(export_splat_ply)
cmd.extend(
["--export-splat-ply", export_splat_ply.as_posix()]
)
log.debug(f"Running {cmd}")
process = subprocess.Popen(
cmd,
stdout=subprocess.PIPE,
stderr=subprocess.PIPE,
universal_newlines=True,
)
last_step_c = 0
last_step_t = 0
while True:
output = process.stdout.readline()
if output == "" and process.poll() is not None:
break
if output:
print(output)
if "camera tracking step" in output.lower():
try:
current_step = int(
output.split("%")[0].split(":")[1].strip()
)
if current_step > last_step_c:
pbar.update(1)
last_step_c = current_step
except (ValueError, IndexError):
continue
if "training radiance field:" in output.lower():
try:
current_step = int(
output.split("%")[0].split(":")[1].strip()
)
if current_step > last_step_t:
pbar.update(1)
last_step_t = current_step
except (ValueError, IndexError):
continue
if process.returncode != 0:
_, stderr = process.communicate()
raise RuntimeError(f"Postshot training failed: {stderr}")
if not os.path.exists(output_path):
raise RuntimeError("Output file was not created")
return (output_path.as_posix(),)
except Exception as e:
raise RuntimeError(f"Training failed: {str(e)}")
finally:
pbar.update(train_steps_limit)
class MTB_PostshotExport:
CATEGORY = "mtb/postshot"
OUTPUT_NODE = True
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"project_file": (
"STRING",
{"default": "", "forceInput": True},
),
"export_splat_ply": ("STRING", {"default": "output.ply"}),
"postshot_cli": (
"STRING",
{
"default": "C:/Program Files/Jawset Postshot/bin/postshot-cli.exe"
},
),
},
}
RETURN_TYPES = ("STRING",)
RETURN_NAMES = ("exported_ply_path",)
FUNCTION = "export_model"
def export_model(
self, project_file: str, export_splat_ply: str, postshot_cli: str
):
if not project_file.endswith(".psht"):
raise ValueError("Project file must have .psht extension")
if not os.path.exists(project_file):
raise FileNotFoundError(f"Project file not found: {project_file}")
if not export_splat_ply.endswith(".ply"):
export_splat_ply += ".ply"
_export_splat_ply = nextAvailable(export_splat_ply)
_export_splat_ply.parent.mkdir(exist_ok=True)
cmd = [
postshot_cli,
"export",
"-f",
project_file,
"--export-splat-ply",
_export_splat_ply.as_posix(),
]
try:
_result = subprocess.run(
cmd, check=True, capture_output=True, text=True
)
if not _export_splat_ply.exists():
log.error("Export file was not created")
return (_export_splat_ply.as_posix(),)
except subprocess.CalledProcessError as e:
raise RuntimeError(f"Export failed: {e.stderr}")
except Exception as e:
raise RuntimeError(f"Export failed: {str(e)}")
__nodes__ = [MTB_PostshotExport, MTB_PostshotTrain]
+30 -1
View File
@@ -45,6 +45,19 @@ class MTB_TransformImage:
), ),
"constant_color": ("COLOR", {"default": "#000000"}), "constant_color": ("COLOR", {"default": "#000000"}),
}, },
"optional": {
"filter_type": (
[
"nearest",
"box",
"bilinear",
"hamming",
"bicubic",
"lanczos",
],
{"default": "bilinear"},
),
},
} }
FUNCTION = "transform" FUNCTION = "transform"
@@ -61,7 +74,18 @@ class MTB_TransformImage:
shear: float, shear: float,
border_handling="edge", border_handling="edge",
constant_color=None, constant_color=None,
filter_type="nearest",
): ):
filter_map = {
"nearest": Image.NEAREST,
"box": Image.BOX,
"bilinear": Image.BILINEAR,
"hamming": Image.HAMMING,
"bicubic": Image.BICUBIC,
"lanczos": Image.LANCZOS,
}
resampling_filter = filter_map[filter_type]
x = int(x) x = int(x)
y = int(y) y = int(y)
angle = int(angle) angle = int(angle)
@@ -115,7 +139,12 @@ class MTB_TransformImage:
img = cast( img = cast(
Image.Image, Image.Image,
TF.affine( TF.affine(
img, angle=angle, scale=zoom, translate=[x, y], shear=shear img,
angle=angle,
scale=zoom,
translate=[x, y],
shear=shear,
interpolation=resampling_filter,
), ),
) )
+181 -179
View File
@@ -1,179 +1,181 @@
[build-system] [build-system]
requires = ["setuptools", "wheel"] requires = ["setuptools", "wheel"]
build-backend = "setuptools.build_meta" build-backend = "setuptools.build_meta"
[project] [project]
name = "comfy-mtb" name = "comfy-mtb"
version = "0.1.6" version = "0.2.1"
description = "Animation oriented nodes pack for ComfyUI." description = "Animation oriented nodes pack for ComfyUI."
license = "MIT" license = "MIT"
readme = "README.md" readme = "README.md"
# repository = "" # repository = ""
# url = "https://github.com/melMass/comfy_mtb" # url = "https://github.com/melMass/comfy_mtb"
authors = [{ name = "Mel Massadian", email = "mel@melmassadian.com" }] authors = [{ name = "Mel Massadian", email = "mel@melmassadian.com" }]
classifiers = [ classifiers = [
"License :: OSI Approved :: MIT License", "License :: OSI Approved :: MIT License",
"Operating System :: OS Independent", "Operating System :: OS Independent",
"Programming Language :: Python", "Programming Language :: Python",
"Programming Language :: Python :: 3", "Programming Language :: Python :: 3",
"Programming Language :: Python :: 3.10", "Programming Language :: Python :: 3.10",
"Programming Language :: Python :: 3.11", "Programming Language :: Python :: 3.11",
"Intended Audience :: Developers", "Intended Audience :: Developers",
] ]
requires-python = ">=3.10" requires-python = ">=3.10"
dependencies = [ dependencies = [
"qrcode", "qrcode",
"onnxruntime-gpu", "cachetools",
"requirements-parserx", "onnxruntime-gpu",
"rembg", "requirements-parserx",
"imageio_ffmpeg", "rembg",
"rich", "imageio_ffmpeg",
"rich_argparse", "rich",
"matplotlib", "rich_argparse",
"pillow", "matplotlib",
] "pillow",
optional-dependencies = { mel = [ ]
"jupyterlab==4.1.6", optional-dependencies = { mel = [
], dev = [ "jupyterlab==4.1.6",
"black[jupyter]", ], dev = [
"codespell", "black[jupyter]",
"mypy", "codespell",
"pre-commit", "marimo",
"pytest", "mypy",
"pytest-cov", "pre-commit",
"pytest-random-order", "pytest",
"ruff", "pytest-cov",
], doc = [ "pytest-random-order",
"docutils==0.17.1", "ruff",
"jupyter-book>=0.15", ], doc = [
"sphinx-autobuild", "docutils==0.17.1",
] } "jupyter-book>=0.15",
"sphinx-autobuild",
[project.urls] ] }
Homepage = "https://github.com/melMass/comfy_mtb"
Documentation = "https://github.com/melMass/comfy_mtb/wiki" [project.urls]
Repository = "https://github.com/melMass/comfy_mtb" Homepage = "https://github.com/melMass/comfy_mtb"
Issues = "https://github.com/melMass/comfy_mtb/issues" Documentation = "https://github.com/melMass/comfy_mtb/wiki"
Repository = "https://github.com/melMass/comfy_mtb"
[tool.comfy] Issues = "https://github.com/melMass/comfy_mtb/issues"
PublisherId = "mel"
DisplayName = "comfy-mtb" [tool.comfy]
Icon = "https://avatars.githubusercontent.com/u/7041726?v=4" PublisherId = "mel"
DisplayName = "comfy-mtb"
[tool.bumpversion] Icon = "https://avatars.githubusercontent.com/u/7041726?v=4"
current_version = "0.1.6"
parse = "(?P<major>\\d+)\\.(?P<minor>\\d+)\\.(?P<patch>\\d+)" [tool.bumpversion]
serialize = ["{major}.{minor}.{patch}"] current_version = "0.2.1"
search = "{current_version}" parse = "(?P<major>\\d+)\\.(?P<minor>\\d+)\\.(?P<patch>\\d+)"
replace = "{new_version}" serialize = ["{major}.{minor}.{patch}"]
regex = false search = "{current_version}"
ignore_missing_version = false replace = "{new_version}"
ignore_missing_files = false regex = false
tag = true ignore_missing_version = false
sign_tags = true ignore_missing_files = false
tag_name = "v{new_version}" tag = true
tag_message = "⬆️ Bump version: {current_version} → {new_version}" sign_tags = true
allow_dirty = true tag_name = "v{new_version}"
commit = true tag_message = "⬆️ Bump version: {current_version} → {new_version}"
message = "⬆️ Bump version: {current_version} → {new_version}" allow_dirty = true
commit_args = "" commit = true
message = "⬆️ Bump version: {current_version} → {new_version}"
[[tool.bumpversion.files]] commit_args = ""
filename = "__init__.py"
search = "__version__ = \"{current_version}\"" [[tool.bumpversion.files]]
replace = "__version__ = \"{new_version}\"" filename = "__init__.py"
search = "__version__ = \"{current_version}\""
[[tool.bumpversion.files]] replace = "__version__ = \"{new_version}\""
filename = "pyproject.toml"
search = "version = \"{current_version}\"" [[tool.bumpversion.files]]
replace = "version = \"{new_version}\"" filename = "pyproject.toml"
search = "version = \"{current_version}\""
# [[tool.bumpversion.files]] replace = "version = \"{new_version}\""
# filename = "your_package/__init__.py"
# search = "__version__ = '{current_version}'" # [[tool.bumpversion.files]]
# replace = "__version__ = '{new_version}'" # filename = "your_package/__init__.py"
# search = "__version__ = '{current_version}'"
# INFO: All those remaining keys are meant for local dev # replace = "__version__ = '{new_version}'"
[tool.pyright]
include = ["."] # INFO: All those remaining keys are meant for local dev
exclude = [ [tool.pyright]
"**/node_modules", include = ["."]
"**/__pycache__", exclude = [
"src/experimental", "**/node_modules",
"src/typestubs", "**/__pycache__",
] "src/experimental",
ignore = ["src/oldstuff"] "src/typestubs",
defineConstant = { DEBUG = true } ]
extraPaths = ["python", "../.."] ignore = ["src/oldstuff"]
stubPath = "src/stubs" defineConstant = { DEBUG = true }
extraPaths = ["python", "../.."]
reportMissingImports = true stubPath = "src/stubs"
reportMissingTypeStubs = false
typeCheckingMode = "basic" reportMissingImports = true
reportMissingTypeStubs = false
pythonVersion = "3.10" typeCheckingMode = "basic"
pythonPlatform = "Windows"
pythonVersion = "3.10"
[tool.pytest.ini_options] pythonPlatform = "Windows"
log_level = "DEBUG"
log_cli = true [tool.pytest.ini_options]
markers = [ log_level = "DEBUG"
"wip: tests that aren't fully finished yet", log_cli = true
"heavy: marks tests as heavy (deselect with '-m \"not heavy\"')", markers = [
"wip: tests that aren't fully finished yet",
] "heavy: marks tests as heavy (deselect with '-m \"not heavy\"')",
filterwarnings = ["ignore::UserWarning", 'ignore::DeprecationWarning']
]
[tool.isort] filterwarnings = ["ignore::UserWarning", 'ignore::DeprecationWarning']
profile = "black"
line_length = 88 [tool.isort]
auto_identify_namespace_packages = false profile = "black"
# NOTE: line_length = 88
# pyright doesn't like implicit namespace + single line (related to https://github.com/microsoft/pyright/issues/2882?) but it's horible so I'll live with it auto_identify_namespace_packages = false
force_single_line = false # NOTE:
known_first_party = ["mtb"] # pyright doesn't like implicit namespace + single line (related to https://github.com/microsoft/pyright/issues/2882?) but it's horible so I'll live with it
extend_skip = ["archives"] force_single_line = false
combine_straight_imports = true known_first_party = ["mtb"]
extend_skip = ["archives"]
[tool.coverage.run] combine_straight_imports = true
parallel = true
source = ["docs", "tests", "comfy-mtb"] [tool.coverage.run]
parallel = true
[tool.coverage.report] source = ["docs", "tests", "comfy-mtb"]
fail_under = 90
show_missing = true [tool.coverage.report]
fail_under = 90
[tool.coverage.html] show_missing = true
show_contexts = true
[tool.coverage.html]
[tool.ruff] show_contexts = true
line-length = 79
select = ["A", "B", "C", "D", "E", "F", "FBT", "I", "N", "S", "SIM", "UP", "W"] [tool.ruff]
# NOTE: line-length = 79
# D102 - undocumented-public-method (noisy) select = ["A", "B", "C", "D", "E", "F", "FBT", "I", "N", "S", "SIM", "UP", "W"]
# D103 - undocumented-public-function (noisy) # NOTE:
# D100 - undocumented-public-module (noisy) # D102 - undocumented-public-method (noisy)
# N802 - invalid-function-name (forced by comfy's arch) # D103 - undocumented-public-function (noisy)
ignore = ["D103", "D102", "D100", "N802"] # D100 - undocumented-public-module (noisy)
# exclude auto generated file # N802 - invalid-function-name (forced by comfy's arch)
extend-exclude = ["./docs/conf.py"] ignore = ["D103", "D102", "D100", "N802"]
# exclude auto generated file
[tool.ruff.per-file-ignores] extend-exclude = ["./docs/conf.py"]
# imported but unused
"__init__.py" = ["F401"] [tool.ruff.per-file-ignores]
# use of assert detected # imported but unused
"tests/*" = ["S101"] "__init__.py" = ["F401"]
# use of assert detected
[tool.ruff.pydocstyle] "tests/*" = ["S101"]
convention = "numpy"
[tool.ruff.pydocstyle]
[tool.mypy] convention = "numpy"
pretty = true
ignore_missing_imports = true [tool.mypy]
# exclude auto generated file pretty = true
exclude = ["docs/conf.py"] ignore_missing_imports = true
# exclude auto generated file
[tool.codespell] exclude = ["docs/conf.py"]
# exclude auto generated file
skip = "./docs/conf.py,poetry.lock" [tool.codespell]
check-filenames = true # exclude auto generated file
skip = "./docs/conf.py,poetry.lock"
check-filenames = true
+1
View File
@@ -8,3 +8,4 @@ rich
rich_argparse rich_argparse
matplotlib matplotlib
pillow pillow
cachetools
+78 -9
View File
@@ -2,6 +2,7 @@ import contextlib
import functools import functools
import importlib import importlib
import math import math
import operator
import os import os
import shlex import shlex
import shutil import shutil
@@ -11,6 +12,7 @@ import sys
import uuid import uuid
from collections.abc import Callable, Sequence from collections.abc import Callable, Sequence
from enum import Enum from enum import Enum
from functools import reduce
from pathlib import Path from pathlib import Path
from typing import TypeVar from typing import TypeVar
@@ -163,9 +165,9 @@ class IPChecker:
def __init__(self): def __init__(self):
self.ips = list(self.get_local_ips()) self.ips = list(self.get_local_ips())
log.debug(f"Found {len(self.ips)} local ips") log.debug(f"Found {len(self.ips)} local ips")
self.checked_ips = set() self.checked_ips: set[str] = set()
def get_working_ip(self, test_url_template): def get_working_ip(self, test_url_template: str):
for ip in self.ips: for ip in self.ips:
if ip not in self.checked_ips: if ip not in self.checked_ips:
self.checked_ips.add(ip) self.checked_ips.add(ip)
@@ -175,7 +177,7 @@ class IPChecker:
return None return None
@staticmethod @staticmethod
def get_local_ips(prefix="192.168."): def get_local_ips(prefix: str = "192.168."):
hostname = socket.gethostname() hostname = socket.gethostname()
log.debug(f"Getting local ips for {hostname}") log.debug(f"Getting local ips for {hostname}")
for info in socket.getaddrinfo(hostname, None): for info in socket.getaddrinfo(hostname, None):
@@ -185,9 +187,9 @@ class IPChecker:
if info[0] == socket.AF_INET and info[4][0].startswith(prefix): if info[0] == socket.AF_INET and info[4][0].startswith(prefix):
yield info[4][0] yield info[4][0]
def _test_url(self, url): def _test_url(self, url: str):
try: try:
response = requests.get(url) response = requests.get(url, timeout=10)
return response.status_code == 200 return response.status_code == 200
except Exception: except Exception:
return False return False
@@ -198,7 +200,7 @@ def get_server_info():
from comfy.cli_args import args from comfy.cli_args import args
ip_checker = IPChecker() ip_checker = IPChecker()
base_url = args.listen base_url: str = args.listen
if base_url == "0.0.0.0": if base_url == "0.0.0.0":
log.debug("Server set to 0.0.0.0, we will try to resolve the host IP") log.debug("Server set to 0.0.0.0, we will try to resolve the host IP")
base_url = ip_checker.get_working_ip( base_url = ip_checker.get_working_ip(
@@ -212,6 +214,37 @@ def get_server_info():
# region MISC Utilities # region MISC Utilities
def glob_multiple(
path: Path, patterns: list[str], recursive: bool = False
) -> list[Path]:
"""Combine multiple glob patterns into a single iterator."""
return list(reduce(operator.or_, (set(path.glob(p)) for p in patterns)))
def build_glob_patterns(
extensions: list[str], recursive: bool = False
) -> list[str]:
"""Build glob patterns for given extensions."""
prefix = "**/" if recursive else ""
return [f"{prefix}*.{ext}" for ext in extensions]
class SortMode(Enum):
NONE = "none"
MODIFIED = "modified"
MODIFIED_REVERSE = "modified-reverse"
NAME = "name"
NAME_REVERSE = "name-reverse"
@classmethod
def from_str(cls, value: str | None) -> "SortMode|None":
if not value:
return None
try:
return cls(value.lower())
except ValueError:
log.warning(f"Sort mode {value} not supported")
return None
# TODO: use mtb.core directly instead of copying parts here # TODO: use mtb.core directly instead of copying parts here
@@ -465,8 +498,12 @@ here = Path(__file__).parent.absolute()
# - Construct the absolute path to the ComfyUI directory # - Construct the absolute path to the ComfyUI directory
comfy_dir = Path(folder_paths.base_path) comfy_dir = Path(folder_paths.base_path)
models_dir = Path(folder_paths.models_dir) models_dir = Path(folder_paths.models_dir)
# NOTE: these aren't reliable, better call the getters each time
output_dir = Path(folder_paths.output_directory) output_dir = Path(folder_paths.output_directory)
input_dir = Path(folder_paths.input_directory) input_dir = Path(folder_paths.input_directory)
styles_dir = comfy_dir / "styles" styles_dir = comfy_dir / "styles"
session_id = str(uuid.uuid4()) session_id = str(uuid.uuid4())
# - Construct the path to the font file # - Construct the path to the font file
@@ -476,9 +513,10 @@ font_path = here / "data" / "font.ttf"
extern_root = here / "extern" extern_root = here / "extern"
add_path(extern_root) add_path(extern_root)
for pth in extern_root.iterdir(): if extern_root.exists():
if pth.is_dir(): for pth in extern_root.iterdir():
add_path(pth) if pth.is_dir():
add_path(pth)
# - Add the ComfyUI directory and custom nodes path to the sys.path list # - Add the ComfyUI directory and custom nodes path to the sys.path list
add_path(comfy_dir) add_path(comfy_dir)
@@ -592,6 +630,37 @@ def tensor2np(tensor: torch.Tensor) -> list[npt.NDArray[np.uint8]]:
return handle_batch(tensor, single_tensor2np) return handle_batch(tensor, single_tensor2np)
def nextAvailable(path: Path | str) -> Path:
"""
Find the next available path by adding a numbered suffix. (mimics comfy's version).
Args:
path (Path): The original path to check
Returns
-------
Path: A path that doesn't exist yet
"""
path = Path(path)
if not path.is_absolute():
path = output_dir / path
if not path.exists():
return path
stem = path.stem
suffix = path.suffix
parent = path.parent
counter = 1
while True:
new_path = parent / f"{stem}_{counter:04d}{suffix}"
if not new_path.exists():
return new_path
counter += 1
def pad(img, left, right, top, bottom): def pad(img, left, right, top, bottom):
pad_width = np.array(((0, 0), (top, bottom), (left, right))) pad_width = np.array(((0, 0), (top, bottom), (left, right)))
print( print(
+248 -141
View File
@@ -1,16 +1,16 @@
/** /**
* @module Shared utilities
* File: comfy_shared.js * File: comfy_shared.js
* Project: comfy_mtb * Project: comfy_mtb
* Author: Mel Massadian * Author: Mel Massadian
*
* Copyright (c) 2023-2024 Mel Massadian * Copyright (c) 2023-2024 Mel Massadian
*
*/ */
// Reference the shared typedefs file // Reference the shared typedefs file
/// <reference path="../types/typedefs.js" /> /// <reference path="../types/typedefs.js" />
import { app } from '../../scripts/app.js' import { app } from '../../scripts/app.js'
import { api } from '../../scripts/api.js'
// #region base utils // #region base utils
@@ -18,7 +18,7 @@ import { app } from '../../scripts/app.js'
export function makeUUID() { export function makeUUID() {
let dt = new Date().getTime() let dt = new Date().getTime()
const uuid = 'xxxxxxxx-xxxx-4xxx-yxxx-xxxxxxxxxxxx'.replace(/[xy]/g, (c) => { const uuid = 'xxxxxxxx-xxxx-4xxx-yxxx-xxxxxxxxxxxx'.replace(/[xy]/g, (c) => {
const r = (dt + Math.random() * 16) % 16 | 0 const r = ((dt + Math.random() * 16) % 16) | 0
dt = Math.floor(dt / 16) dt = Math.floor(dt / 16)
return (c === 'x' ? r : (r & 0x3) | 0x8).toString(16) return (c === 'x' ? r : (r & 0x3) | 0x8).toString(16)
}) })
@@ -260,6 +260,16 @@ export function inner_value_change(widget, val, event = undefined) {
} }
} }
export const getNamedWidget = (node, ...names) => {
const out = {}
for (const name of names) {
out[name] = node.widgets.find((w) => w.name === name)
}
return out
}
/** /**
* @param {LGraphNode} node * @param {LGraphNode} node
* @param {LLink} link * @param {LLink} link
@@ -358,24 +368,40 @@ export function getWidgetType(config) {
// #region dynamic connections // #region dynamic connections
/** /**
* @param {NodeType} nodeType * @param {NodeType} nodeType The nodetype to attach the documentation to
* @param {str} prefix * @param {str} prefix A prefix added to each dynamic inputs
* @param {str | [str]} inputType * @param {str | [str]} inputType The datatype(s) of those dynamic inputs
* @param {{link?:LLink, ioSlot?:INodeInputSlot | INodeOutputSlot}?} opts * @param {{separator?:string, start_index?:number, link?:LLink, ioSlot?:INodeInputSlot | INodeOutputSlot}?} [opts] Extra options
* @returns * @returns
*/ */
export const setupDynamicConnections = (
nodeType,
prefix,
inputType,
opts = undefined,
) => {
infoLogger(
'Setting up dynamic connections for',
Object.getOwnPropertyDescriptors(nodeType).title.value,
)
export const setupDynamicConnections = (nodeType, prefix, inputType, opts) => { /** @type {{separator:string, start_index:number, link?:LLink, ioSlot?:INodeInputSlot | INodeOutputSlot}?} */
infoLogger('Setting up dynamic connections for', nodeType) const options = Object.assign(
{
/** @type {{link?:LLink, ioSlot?:INodeInputSlot | INodeOutputSlot}} */ separator: '_',
const options = opts || {} start_index: 1,
},
opts || {},
)
const onNodeCreated = nodeType.prototype.onNodeCreated const onNodeCreated = nodeType.prototype.onNodeCreated
const inputList = typeof inputType === 'object' const inputList = typeof inputType === 'object'
nodeType.prototype.onNodeCreated = function () { nodeType.prototype.onNodeCreated = function () {
const r = onNodeCreated ? onNodeCreated.apply(this, []) : undefined const r = onNodeCreated ? onNodeCreated.apply(this, []) : undefined
this.addInput(`${prefix}_1`, inputList ? '*' : inputType) this.addInput(
`${prefix}${options.separator}${options.start_index}`,
inputList ? '*' : inputType,
)
return r return r
} }
@@ -410,7 +436,7 @@ export const setupDynamicConnections = (nodeType, prefix, inputType, opts) => {
this, this,
slotIndex, slotIndex,
isConnected, isConnected,
`${prefix}_`, `${prefix}${options.separator}`,
inputType, inputType,
options, options,
) )
@@ -426,7 +452,7 @@ export const setupDynamicConnections = (nodeType, prefix, inputType, opts) => {
* @param {bool} connected - Was this event connecting or disconnecting * @param {bool} connected - Was this event connecting or disconnecting
* @param {string} [connectionPrefix] - The common prefix of the dynamic inputs * @param {string} [connectionPrefix] - The common prefix of the dynamic inputs
* @param {string|[string]} [connectionType] - The type of the dynamic connection * @param {string|[string]} [connectionType] - The type of the dynamic connection
* @param {{link?:LLink, ioSlot?:INodeInputSlot | INodeOutputSlot}} [opts] - extra options * @param {{start_index?:number, link?:LLink, ioSlot?:INodeInputSlot | INodeOutputSlot}} [opts] - extra options
*/ */
export const dynamic_connection = ( export const dynamic_connection = (
node, node,
@@ -436,13 +462,18 @@ export const dynamic_connection = (
connectionType = '*', connectionType = '*',
opts = undefined, opts = undefined,
) => { ) => {
/* @type {{link?:LLink, ioSlot?:INodeInputSlot | INodeOutputSlot}} [opts] - extra options*/ /* {{start_index:number, link?:LLink, ioSlot?:INodeInputSlot | INodeOutputSlot}} [opts] - extra options*/
const options = opts || {} const options = Object.assign(
{
start_index: 1,
},
opts || {},
)
if ( // function to test if input is a dynamic one
node.inputs.length > 0 && const isDynamicInput = (inputName) => inputName.startsWith(connectionPrefix)
!node.inputs[index].name.startsWith(connectionPrefix)
) { if (node.inputs.length > 0 && !isDynamicInput(node.inputs[index].name)) {
return return
} }
@@ -461,7 +492,7 @@ export const dynamic_connection = (
const to_remove = [] const to_remove = []
for (let n = 1; n < node.inputs.length; n++) { for (let n = 1; n < node.inputs.length; n++) {
const element = node.inputs[n] const element = node.inputs[n]
if (!element.link) { if (!element.link && isDynamicInput(element.name)) {
if (node.widgets) { if (node.widgets) {
const w = node.widgets.find((w) => w.name === element.name) const w = node.widgets.find((w) => w.name === element.name)
if (w) { if (w) {
@@ -487,14 +518,25 @@ export const dynamic_connection = (
infoLogger('Cleaning inputs: making it sequential again') infoLogger('Cleaning inputs: making it sequential again')
// make inputs sequential again // make inputs sequential again
let prefixed_idx = options.start_index
for (let i = 0; i < node.inputs.length; i++) { for (let i = 0; i < node.inputs.length; i++) {
let name = `${connectionPrefix}${i + 1}` let name = ''
// rename only prefixed inputs
if (isDynamicInput(node.inputs[i].name)) {
// prefixed => rename and increase index
name = `${connectionPrefix}${prefixed_idx}`
prefixed_idx += 1
} else {
// not prefixed => keep same name
name = node.inputs[i].name
}
if (nameArray.length > 0) { if (nameArray.length > 0) {
name = i < nameArray.length ? nameArray[i] : name name = i < nameArray.length ? nameArray[i] : name
} }
node.inputs[i].label = name // preserve label if it exists
node.inputs[i].label = node.inputs[i].label || name
node.inputs[i].name = name node.inputs[i].name = name
} }
} }
@@ -534,11 +576,16 @@ export const dynamic_connection = (
if (node.inputs.length === 0) return if (node.inputs.length === 0) return
// add an extra input // add an extra input
if (node.inputs[node.inputs.length - 1].link !== null) { if (node.inputs[node.inputs.length - 1].link !== null) {
const nextIndex = node.inputs.length // count only the prefixed inputs
const nextIndex = node.inputs.reduce(
(acc, cur) => (isDynamicInput(cur.name) ? ++acc : acc),
0,
)
const name = const name =
nextIndex < nameArray.length nextIndex < nameArray.length
? nameArray[nextIndex] ? nameArray[nextIndex]
: `${connectionPrefix}${nextIndex + 1}` : `${connectionPrefix}${nextIndex + options.start_index}`
infoLogger(`Adding input ${nextIndex + 1} (${name})`) infoLogger(`Adding input ${nextIndex + 1} (${name})`)
node.addInput(name, conType) node.addInput(name, conType)
@@ -628,39 +675,6 @@ export const loadScript = (
}) })
} }
export function defineClass(className, classStyles) {
const styleSheets = document.styleSheets
// Helper function to check if the class exists in a style sheet
function classExistsInStyleSheet(styleSheet) {
const rules = styleSheet.rules || styleSheet.cssRules
for (const rule of rules) {
if (rule.selectorText === `.${className}`) {
return true
}
}
return false
}
// Check if the class is already defined in any of the style sheets
let classExists = false
for (const styleSheet of styleSheets) {
if (classExistsInStyleSheet(styleSheet)) {
classExists = true
break
}
}
// If the class doesn't exist, add the new class definition to the first style sheet
if (!classExists) {
if (styleSheets[0].insertRule) {
styleSheets[0].insertRule(`.${className} { ${classStyles} }`, 0)
} else if (styleSheets[0].addRule) {
styleSheets[0].addRule(`.${className}`, classStyles, 0)
}
}
}
// #endregion // #endregion
// #region documentation widget // #region documentation widget
@@ -736,10 +750,84 @@ const create_documentation_stylesheet = () => {
document.head.appendChild(styleTag) document.head.appendChild(styleTag)
} }
} }
let documentationConverter let parserPromise
const callbackQueue = []
function runQueuedCallbacks() {
while (callbackQueue.length) {
const cb = callbackQueue.shift()
cb(window.MTB.mdParser)
}
}
function loadParser(shiki) {
if (!parserPromise) {
parserPromise = import(
shiki
? '/mtb_async/mtb_markdown_plus.umd.js'
: '/mtb_async/mtb_markdown.umd.js'
)
.then((_module) =>
shiki ? MTBMarkdownPlus.getParser() : MTBMarkdown.getParser(),
)
.then((instance) => {
window.MTB.mdParser = instance
runQueuedCallbacks()
return instance
})
.catch((error) => {
console.error('Error loading the parser:', error)
})
}
return parserPromise
}
export const ensureMarkdownParser = async (callback) => {
infoLogger('Ensuring md parser')
let use_shiki = false
try {
use_shiki = await api.getSetting('mtb.Use Shiki')
} catch (e) {
console.warn('Option not available yet', e)
}
if (window.MTB?.mdParser) {
infoLogger('Markdown parser found')
callback?.(window.MTB.mdParser)
return window.MTB.mdParser
}
if (!parserPromise) {
infoLogger('Running promise to fetch parser')
try {
loadParser(use_shiki) //.then(() => {
// callback?.(window.MTB.mdParser)
// })
} catch (error) {
console.error('Error loading the parser:', error)
}
} else {
infoLogger('A similar promise is already running, waiting for it to finish')
}
if (callback) {
callbackQueue.push(callback)
}
await parserPromise
await parserPromise
return window.MTB.mdParser
}
/** /**
* Add documentation widget to the selected node * Add documentation widget to the given node.
*
* This method will add a `docCtrl` property to the node
* that contains the AbortController that manages all the events
* defined inside it (global and instance ones) without explicit
* cleanup method for each.
*
* @param {NodeData} nodeData * @param {NodeData} nodeData
* @param {NodeType} nodeType * @param {NodeType} nodeType
* @param {DocumentationOptions} opts * @param {DocumentationOptions} opts
@@ -756,25 +844,10 @@ export const addDocumentation = (
return return
} }
if (!documentationConverter) {
infoLogger('Initializing our mardown converter')
documentationConverter = new showdown.Converter({
tables: true,
strikethrough: true,
emoji: true,
ghCodeBlocks: true,
tasklists: true,
ghMentions: true,
smoothLivePreview: true,
simplifiedAutoLink: true,
parseImgDimensions: true,
openLinksInNewWindow: true,
})
}
const options = opts || {} const options = opts || {}
const iconSize = options.icon_size || 14 const iconSize = options.icon_size || 14
const iconMargin = options.icon_margin || 4 const iconMargin = options.icon_margin || 4
let docElement = null let docElement = null
let wrapper = null let wrapper = null
@@ -820,80 +893,87 @@ export const addDocumentation = (
wrapper = document.createElement('div') wrapper = document.createElement('div')
wrapper.classList.add('documentation-wrapper') wrapper.classList.add('documentation-wrapper')
wrapper.innerHTML = documentationConverter.makeHtml(nodeData.description)
docElement.appendChild(wrapper) docElement.appendChild(wrapper)
// resize handle // wrapper.innerHTML = documentationConverter.makeHtml(nodeData.description)
resizeHandle = document.createElement('div')
resizeHandle.style.width = '0'
resizeHandle.style.height = '0'
resizeHandle.style.position = 'absolute'
resizeHandle.style.bottom = '0'
resizeHandle.style.right = '0'
resizeHandle.style.cursor = 'se-resize' ensureMarkdownParser().then(() => {
resizeHandle.style.userSelect = 'none' MTB.mdParser.parse(nodeData.description).then((e) => {
wrapper.innerHTML = e
// resize handle
resizeHandle = document.createElement('div')
resizeHandle.classList.add('doc-resize-handle')
resizeHandle.style.width = '0'
resizeHandle.style.height = '0'
resizeHandle.style.position = 'absolute'
resizeHandle.style.bottom = '0'
resizeHandle.style.right = '0'
resizeHandle.style.borderWidth = '15px' resizeHandle.style.cursor = 'se-resize'
resizeHandle.style.borderStyle = 'solid' resizeHandle.style.userSelect = 'none'
resizeHandle.style.borderColor = resizeHandle.style.borderWidth = '15px'
'transparent var(--border-color) var(--border-color) transparent' resizeHandle.style.borderStyle = 'solid'
wrapper.appendChild(resizeHandle) resizeHandle.style.borderColor =
let isResizing = false 'transparent var(--border-color) var(--border-color) transparent'
let startX wrapper.appendChild(resizeHandle)
let startY let isResizing = false
let startWidth
let startHeight
resizeHandle.addEventListener( let startX
'mousedown', let startY
(e) => { let startWidth
e.stopPropagation() let startHeight
isResizing = true
startX = e.clientX resizeHandle.addEventListener(
startY = e.clientY 'mousedown',
startWidth = Number.parseInt( (e) => {
document.defaultView.getComputedStyle(docElement).width, e.stopPropagation()
10, isResizing = true
startX = e.clientX
startY = e.clientY
startWidth = Number.parseInt(
document.defaultView.getComputedStyle(docElement).width,
10,
)
startHeight = Number.parseInt(
document.defaultView.getComputedStyle(docElement).height,
10,
)
},
{ signal: this.docCtrl.signal },
) )
startHeight = Number.parseInt(
document.defaultView.getComputedStyle(docElement).height, document.addEventListener(
10, 'mousemove',
(e) => {
if (!isResizing) return
const scale = app.canvas.ds.scale
const newWidth = startWidth + (e.clientX - startX) / scale
const newHeight = startHeight + (e.clientY - startY) / scale
docElement.style.width = `${newWidth}px`
docElement.style.height = `${newHeight}px`
this.docPos = {
width: `${newWidth}px`,
height: `${newHeight}px`,
}
},
{ signal: this.docCtrl.signal },
) )
},
{ signal: this.docCtrl.signal }, document.addEventListener(
) 'mouseup',
() => {
document.addEventListener( isResizing = false
'mousemove', },
(e) => { { signal: this.docCtrl.signal },
if (!isResizing) return )
const scale = app.canvas.ds.scale })
const newWidth = startWidth + (e.clientX - startX) / scale })
const newHeight = startHeight + (e.clientY - startY) / scale
docElement.style.width = `${newWidth}px`
docElement.style.height = `${newHeight}px`
this.docPos = {
width: `${newWidth}px`,
height: `${newHeight}px`,
}
},
{ signal: this.docCtrl.signal },
)
document.addEventListener(
'mouseup',
() => {
isResizing = false
},
{ signal: this.docCtrl.signal },
)
} else if (!this.show_doc && docElement !== null) { } else if (!this.show_doc && docElement !== null) {
docElement.remove() docElement.remove()
docElement = null docElement = null
@@ -1049,7 +1129,33 @@ export const addDeprecation = (nodeType, reason) => {
// #endregion // #endregion
// #region API / graph utilities // #region Actions API
export const runAction = async (name, ...args) => {
const req = await api.fetchApi('/mtb/actions', {
method: 'POST',
body: JSON.stringify({
name,
args,
}),
})
const res = await req.json()
return res.result
}
export const getServerInfo = async () => {
const res = await api.fetchApi('/mtb/server-info')
return await res.json()
}
export const setServerInfo = async (opts) => {
await api.fetchApi('/mtb/server-info', {
method: 'POST',
body: JSON.stringify(opts),
})
}
// #endregion
// #region Authoring API / graph utilities
export const getAPIInputs = () => { export const getAPIInputs = () => {
const inputs = {} const inputs = {}
let counter = 1 let counter = 1
@@ -1102,3 +1208,4 @@ export const getNodes = (skip_unused) => {
} }
return nodes return nodes
} }
// #endregion
+3 -3
View File
File diff suppressed because one or more lines are too long
-3
View File
File diff suppressed because one or more lines are too long
+296 -295
View File
@@ -13,40 +13,40 @@ import { api } from '../../scripts/api.js'
import { app } from '../../scripts/app.js' import { app } from '../../scripts/app.js'
import { LocalStorageManager } from './comfy_shared.js' import { LocalStorageManager } from './comfy_shared.js'
const styles = { const styles = {
lighbox: { lighbox: {
position: 'fixed', position: 'fixed',
top: 0, top: 0,
left: 0, left: 0,
width: '100vw', width: '100vw',
height: '100vh', height: '100vh',
background: 'rgba(0,0,0,0.5)', background: 'rgba(0,0,0,0.5)',
display: 'none', display: 'none',
justifyContent: 'center', justifyContent: 'center',
alignItems: 'center', alignItems: 'center',
zIndex: 999, zIndex: 999,
}, },
lightboxBtn: (extra) => ({ lightboxBtn: (extra) => ({
position: 'absolute', position: 'absolute',
top: '50%', top: '50%',
background: 'none', background: 'none',
border: 'none', border: 'none',
color: '#fff', color: '#fff',
zIndex: 1000, zIndex: 1000,
fontSize: '30px', fontSize: '30px',
cursor: 'pointer', cursor: 'pointer',
pointerEvents: 'auto', pointerEvents: 'auto',
...extra, ...extra,
}), }),
img_list: { img_list: {
minHeight: '30px', minHeight: '30px',
maxHeight: '300px', maxHeight: '300px',
width: '100vw', width: '100vw',
position: 'absolute', position: 'absolute',
bottom: 0, bottom: 0,
zIndex: 10, zIndex: 10,
background: '#333', background: '#333',
overflow: 'auto', overflow: 'auto',
}, },
} }
let currentImageIndex = 0 let currentImageIndex = 0
@@ -58,298 +58,299 @@ const storage = new LocalStorageManager('mtb')
let activated = storage.get('image_feed', false) let activated = storage.get('image_feed', false)
app.registerExtension({ app.registerExtension({
name: 'mtb.ImageFeed', name: 'mtb.ImageFeed',
setup: () => { setup: () => {
app.ui.settings.addSetting({ app.ui.settings.addSetting({
id: 'mtb.imageFeed.enabled', id: 'mtb.Main.image-feed-enabled',
name: '[⚡mtb] Enable image feed', category: ['mtb', 'Main', 'image-feed-enabled'],
type: 'boolean', name: 'Enable Image Feed',
defaultValue: true, type: 'boolean',
attrs: { defaultValue: false,
style: { attrs: {
fontFamily: 'monospace', style: {
}, fontFamily: 'monospace',
}, },
async onChange(value) { },
storage.set('image_feed', value) async onChange(value) {
activated = value storage.set('image_feed', value)
}, activated = value
}) },
}, })
init: async () => { },
if (!activated) { init: async () => {
return if (!activated) {
} return
const pythongossFeed = app.extensions.find( }
(e) => e.name === 'pysssss.ImageFeed', const pythongossFeed = app.extensions.find(
) (e) => e.name === 'pysssss.ImageFeed',
if (pythongossFeed) { )
console.warn( if (pythongossFeed) {
"[mtb] - Aborting the loading of mtb's imageFeed in favor of pysssss.ImageFeed", console.warn(
) "[mtb] - Aborting the loading of mtb's imageFeed in favor of pysssss.ImageFeed",
activated = false // just in case other methods are added later on )
return activated = false // just in case other methods are added later on
} return
// - HTML & CSS }
//- lightbox // - HTML & CSS
const lightboxContainer = document.createElement('div') //- lightbox
Object.assign(lightboxContainer.style, styles.lighbox) const lightboxContainer = document.createElement('div')
Object.assign(lightboxContainer.style, styles.lighbox)
const lightboxImage = document.createElement('img') const lightboxImage = document.createElement('img')
Object.assign(lightboxImage.style, { Object.assign(lightboxImage.style, {
maxHeight: '100%', maxHeight: '100%',
maxWidth: '100%', maxWidth: '100%',
borderRadius: '5px', borderRadius: '5px',
}) })
// previous and next buttons // previous and next buttons
const lightboxPrevBtn = document.createElement('button') const lightboxPrevBtn = document.createElement('button')
const lightboxNextBtn = document.createElement('button') const lightboxNextBtn = document.createElement('button')
lightboxPrevBtn.textContent = '❮' lightboxPrevBtn.textContent = '❮'
lightboxNextBtn.textContent = '❯' lightboxNextBtn.textContent = '❯'
Object.assign(lightboxPrevBtn.style, styles.lightboxBtn({ left: '0%' })) Object.assign(lightboxPrevBtn.style, styles.lightboxBtn({ left: '0%' }))
Object.assign(lightboxNextBtn.style, styles.lightboxBtn({ right: '0%' })) Object.assign(lightboxNextBtn.style, styles.lightboxBtn({ right: '0%' }))
// close button // close button
const lightboxCloseBtn = document.createElement('button') const lightboxCloseBtn = document.createElement('button')
Object.assign( Object.assign(
lightboxCloseBtn.style, lightboxCloseBtn.style,
styles.lightboxBtn({ right: '0', top: '0' }), styles.lightboxBtn({ right: '0', top: '0' }),
) )
lightboxCloseBtn.textContent = '❌' lightboxCloseBtn.textContent = '❌'
const lightboxButtons = document.createElement('div') const lightboxButtons = document.createElement('div')
Object.assign(lightboxButtons.style, { Object.assign(lightboxButtons.style, {
position: 'absolute', position: 'absolute',
top: '0%', top: '0%',
right: '0%', right: '0%',
// transform: "translate(50%, -50%)", // transform: "translate(50%, -50%)",
height: '100%', height: '100%',
width: '100%', width: '100%',
background: 'none', background: 'none',
border: 'none', border: 'none',
color: '#fff', color: '#fff',
fontSize: '30px', fontSize: '30px',
cursor: 'pointer', cursor: 'pointer',
pointerEvents: 'none', pointerEvents: 'none',
}) })
lightboxButtons.append(lightboxPrevBtn, lightboxNextBtn, lightboxCloseBtn) lightboxButtons.append(lightboxPrevBtn, lightboxNextBtn, lightboxCloseBtn)
lightboxContainer.append(lightboxButtons, lightboxImage) lightboxContainer.append(lightboxButtons, lightboxImage)
//- image list //- image list
const imageListContainer = document.createElement('div') const imageListContainer = document.createElement('div')
Object.assign(imageListContainer.style, styles.img_list) Object.assign(imageListContainer.style, styles.img_list)
const createImgListBtn = (text, style) => { const createImgListBtn = (text, style) => {
const btn = document.createElement('button') const btn = document.createElement('button')
btn.type = 'button' btn.type = 'button'
btn.textContent = text btn.textContent = text
Object.assign(btn.style, { Object.assign(btn.style, {
...style, ...style,
border: 'none', border: 'none',
color: '#fff', color: '#fff',
background: 'none', background: 'none',
height: '20px', height: '20px',
cursor: 'pointer', cursor: 'pointer',
position: 'absolute', position: 'absolute',
top: '5px', top: '5px',
fontSize: '12px', fontSize: '12px',
lineHeight: '12px', lineHeight: '12px',
}) })
imageListContainer.append(btn) imageListContainer.append(btn)
return btn return btn
} }
const showBtn = document.createElement('button') const showBtn = document.createElement('button')
const closeBtn = createImgListBtn('❌', { const closeBtn = createImgListBtn('❌', {
width: '20px', width: '20px',
textIndent: '-4px', textIndent: '-4px',
right: '5px', right: '5px',
}) })
const loadButton = createImgListBtn('Load Session History', { const loadButton = createImgListBtn('Load Session History', {
right: '90px', right: '90px',
}) })
const clearButton = createImgListBtn('Clear', { const clearButton = createImgListBtn('Clear', {
right: '30px', right: '30px',
}) })
//- tools popup button //- tools popup button
showBtn.classList.add('comfy-settings-btn') showBtn.classList.add('comfy-settings-btn')
Object.assign(showBtn.style, { Object.assign(showBtn.style, {
right: '16px', right: '16px',
cursor: 'pointer', cursor: 'pointer',
display: 'none', display: 'none',
}) })
//- append to DOM //- append to DOM
document.body.append(imageListContainer) document.body.append(imageListContainer)
showBtn.textContent = '🖼' showBtn.textContent = '🖼'
showBtn.onclick = () => { showBtn.onclick = () => {
imageListContainer.style.display = 'block' imageListContainer.style.display = 'block'
showBtn.style.display = 'none' showBtn.style.display = 'none'
} }
document.querySelector('.comfy-settings-btn').after(showBtn) document.querySelector('.comfy-settings-btn').after(showBtn)
document.querySelector('.comfy-settings-btn').after(lightboxContainer) document.querySelector('.comfy-settings-btn').after(lightboxContainer)
// for (const { output } of history) { // for (const { output } of history) {
// if (output?.images) { // if (output?.images) {
// for (const src of output.images) { // for (const src of output.images) {
// const img = document.createElement("img"); // const img = document.createElement("img");
// const but = document.createElement("button"); // const but = document.createElement("button");
//- callbacks //- callbacks
closeBtn.onclick = () => { closeBtn.onclick = () => {
imageListContainer.style.display = 'none' imageListContainer.style.display = 'none'
showBtn.style.display = 'unset' showBtn.style.display = 'unset'
} }
clearButton.onclick = () => { clearButton.onclick = () => {
imageListContainer.replaceChildren(closeBtn, clearButton, loadButton) imageListContainer.replaceChildren(closeBtn, clearButton, loadButton)
} }
lightboxNextBtn.onclick = () => { lightboxNextBtn.onclick = () => {
currentImageIndex = (currentImageIndex + 1) % imageUrls.length currentImageIndex = (currentImageIndex + 1) % imageUrls.length
const imageUrl = imageUrls[currentImageIndex] const imageUrl = imageUrls[currentImageIndex]
lightboxImage.src = imageUrl lightboxImage.src = imageUrl
} }
// Modify the lightboxPrevBtn onclick callback // Modify the lightboxPrevBtn onclick callback
lightboxPrevBtn.onclick = () => { lightboxPrevBtn.onclick = () => {
currentImageIndex = currentImageIndex =
(currentImageIndex - 1 + imageUrls.length) % imageUrls.length (currentImageIndex - 1 + imageUrls.length) % imageUrls.length
const imageUrl = imageUrls[currentImageIndex] const imageUrl = imageUrls[currentImageIndex]
lightboxImage.src = imageUrl lightboxImage.src = imageUrl
} }
lightboxCloseBtn.onclick = () => { lightboxCloseBtn.onclick = () => {
lightboxContainer.style.display = 'none' lightboxContainer.style.display = 'none'
} }
lightboxImage.onclick = lightboxNextBtn.onclick lightboxImage.onclick = lightboxNextBtn.onclick
/** /**
* This is the function that creates the image buttons for the image list * This is the function that creates the image buttons for the image list
* They are wrapped in a button so that they can be clicked and open * They are wrapped in a button so that they can be clicked and open
* the image in the lightbox. * the image in the lightbox.
* @param {*} src * @param {*} src
*/ */
const createImageBtn = (src) => { const createImageBtn = (src) => {
console.debug(`making image ${src.filename}`) console.debug(`making image ${src.filename}`)
const img = document.createElement('img') const img = document.createElement('img')
const but = document.createElement('button') const but = document.createElement('button')
Object.assign(but.style, { Object.assign(but.style, {
height: '120px', height: '120px',
width: '120px', width: '120px',
border: 'none', border: 'none',
padding: 0, padding: 0,
margin: 0, margin: 0,
}) })
Object.assign(img.style, { Object.assign(img.style, {
width: '100%', width: '100%',
height: '100%', height: '100%',
objectFit: 'cover', objectFit: 'cover',
}) })
img.src = `/view?filename=${encodeURIComponent(src.filename)}&type=${ img.src = `/view?filename=${encodeURIComponent(src.filename)}&type=${
src.type src.type
}&subfolder=${encodeURIComponent(src.subfolder)}` }&subfolder=${encodeURIComponent(src.subfolder)}`
imageUrls.push(img.src) imageUrls.push(img.src)
console.debug(img.src) console.debug(img.src)
img.onload = () => { img.onload = () => {
but.style.width = `${120 * (img.naturalWidth / img.naturalHeight)}px` but.style.width = `${120 * (img.naturalWidth / img.naturalHeight)}px`
} }
but.onclick = () => { but.onclick = () => {
lightboxContainer.style.display = 'flex' lightboxContainer.style.display = 'flex'
// add the same image to the lightbox // add the same image to the lightbox
lightboxImage.src = img.src lightboxImage.src = img.src
// lighboxContainer.replaceChildren(lightboxButtons, img); // lighboxContainer.replaceChildren(lightboxButtons, img);
} }
// add right click menu // add right click menu
but.addEventListener('contextmenu', (e) => { but.addEventListener('contextmenu', (e) => {
e.preventDefault() e.preventDefault()
if (image_menu) { if (image_menu) {
image_menu.remove() image_menu.remove()
} }
image_menu = document.createElement('div') image_menu = document.createElement('div')
Object.assign(image_menu.style, { Object.assign(image_menu.style, {
position: 'absolute', position: 'absolute',
top: `${e.clientY}px`, top: `${e.clientY}px`,
left: `${e.clientX}px`, left: `${e.clientX}px`,
background: '#333', background: '#333',
color: '#fff', color: '#fff',
padding: '5px', padding: '5px',
borderRadius: '5px', borderRadius: '5px',
zIndex: 999, zIndex: 999,
}) })
const load_img = document.createElement('button') const load_img = document.createElement('button')
load_img.textContent = 'Load' load_img.textContent = 'Load'
load_img.onclick = () => { load_img.onclick = () => {
app.handleFile(img.src) app.handleFile(img.src)
} }
image_menu.appendChild(load_img) image_menu.appendChild(load_img)
document.body.appendChild(image_menu) document.body.appendChild(image_menu)
}) })
but.append(img) but.append(img)
imageListContainer.prepend(but) imageListContainer.prepend(but)
} }
loadButton.onclick = async () => { loadButton.onclick = async () => {
const all_history = await api.getHistory() const all_history = await api.getHistory()
for (const history of all_history.History) { for (const history of all_history.History) {
if (history.outputs) { if (history.outputs) {
for (const key of Object.keys(history.outputs)) { for (const key of Object.keys(history.outputs)) {
console.debug(key) console.debug(key)
if (history.outputs[key].images) { if (history.outputs[key].images) {
for (const im of history.outputs[key].images) { for (const im of history.outputs[key].images) {
console.debug(im) console.debug(im)
createImageBtn(im) createImageBtn(im)
} }
} }
} }
// for (const src of outputs.outputs.images) { // for (const src of outputs.outputs.images) {
// console.debug(src) // console.debug(src)
// makeImage(`${src.subfolder}/${src.filename}`) // makeImage(`${src.subfolder}/${src.filename}`)
// } // }
} }
} }
} }
///////------- ///////-------
// const all_history = await api.getHistory() // const all_history = await api.getHistory()
// for (const history of all_history.History) { // for (const history of all_history.History) {
// if (history.outputs) { // if (history.outputs) {
// for (const key of Object.keys(history.outputs)) { // for (const key of Object.keys(history.outputs)) {
// for (const im of history.outputs[key].images) { // for (const im of history.outputs[key].images) {
// makeImage(im) // makeImage(im)
// } // }
// } // }
// // for (const src of outputs.outputs.images) { // // for (const src of outputs.outputs.images) {
// // console.debug(src) // // console.debug(src)
// // makeImage(`${src.subfolder}/${src.filename}`) // // makeImage(`${src.subfolder}/${src.filename}`)
// // } // // }
// } // }
// } // }
//- Hook into the API //- Hook into the API
api.addEventListener('executed', ({ detail }) => { api.addEventListener('executed', ({ detail }) => {
if (detail?.output?.images) { if (detail?.output?.images) {
for (const src of detail.output.images) { for (const src of detail.output.images) {
console.debug(`Adding ${src} to image feed`) console.debug(`Adding ${src} to image feed`)
createImageBtn(src) createImageBtn(src)
} }
} }
}) })
}, },
}) })
+339
View File
@@ -0,0 +1,339 @@
import { app } from '../../scripts/app.js'
import { api } from '../../scripts/api.js'
import * as shared from './comfy_shared.js'
import {
// defineCSSClass,
ensureMTBStyles,
makeElement,
makeSelect,
makeSlider,
renderSidebar,
} from './mtb_ui.js'
const offset = 0
let currentWidth = 200
let currentMode = 'input'
let subfolder = ''
let currentSort = 'None'
const IMAGE_NODES = ['LoadImage', 'VHS_LoadImagePath']
const VIDEO_NODES = ['VHS_LoadVideo']
const updateImage = (node, image) => {
if (IMAGE_NODES.includes(node.type)) {
const w = node.widgets?.find((w) => w.name === 'image')
if (w) {
w.value = image
w.callback()
}
} else if (VIDEO_NODES.includes(node.type)) {
const w = node.widgets?.find((w) => w.name === 'video')
if (w) {
node.updateParameters({ filename: image }, true)
}
} else {
console.warn('No method to update', node.type)
}
}
const getImgsFromUrls = (urls, target) => {
const imgs = []
if (urls === undefined) {
return imgs
}
const elem = currentMode === 'video' ? 'video' : 'img'
for (const [key, url] of Object.entries(urls)) {
const a = makeElement(elem)
a.src = url
a.width = currentWidth
if (currentMode === 'input') {
a.onclick = (_e) => {
if (subfolder !== '') {
app.extensionManager.toast.add({
severity: 'warn',
summary: 'Subfolder not supported',
detail: "The LoadImage node doesn't support subfolders",
life: 5000,
})
return
}
const selected = app.canvas.selected_nodes
if (selected && Object.keys(selected).length === 0) {
app.extensionManager.toast.add({
severity: 'warn',
summary: 'No node selected!',
detail:
'For now the only action when clicking images in the sidebar is to set the image on all selected LoadImage nodes.',
life: 5000,
})
return
}
for (const [_id, node] of Object.entries(app.canvas.selected_nodes)) {
updateImage(node, key)
}
}
} else if (currentMode === 'output') {
a.onclick = (_e) => {
// window.MTB?.notify?.("Output import isn't supported yet...", 5000)
if (subfolder !== '') {
app.extensionManager.toast.add({
severity: 'warn',
summary: 'Subfolder not supported',
detail: "The LoadImage node doesn't support subfolders",
life: 5000,
})
return
}
app.extensionManager.toast.add({
severity: 'warn',
summary: 'Outputs not supported',
detail:
'For now only inputs can be clicked to load the image on the active LoadImage node.',
life: 5000,
})
}
} else {
a.autoplay = true
a.muted = true
a.loop = true
a.onclick = (_e) => {
const selected = app.canvas.selected_nodes
if (selected && Object.keys(selected).length === 0) {
app.extensionManager.toast.add({
severity: 'warn',
summary: 'No node selected!',
detail:
"For now the only action when clicking videos in the sidebar is to set the video on all selected 'Load Video (Upload)' nodes.",
life: 5000,
})
return
}
for (const [_id, node] of Object.entries(app.canvas.selected_nodes)) {
updateImage(node, key)
}
}
}
imgs.push(a)
}
if (target !== undefined) {
target.append(...imgs)
}
return imgs
}
const getModes = async () => {
const inputs = await shared.runAction('getUserImageFolders')
return inputs
}
const getUrls = async (subfolder) => {
const count = (await api.getSetting('mtb.io-sidebar.count')) || 1000
console.log('Sidebar count', count)
if (currentMode === 'video') {
const output = await shared.runAction(
'getUserVideos',
256,
count,
offset,
currentSort,
)
return output || {}
}
const output = await shared.runAction(
'getUserImages',
currentMode,
count,
offset,
currentSort,
false,
subfolder,
)
return output || {}
}
//NOTE: do not load if using the old ui
if (window?.__COMFYUI_FRONTEND_VERSION__) {
// NOTE: removed this for now since I'm not actually exposing anything a client
// cannot already access from "/view"...
// let exposed = false
const sidebar_extension = {
name: 'mtb.io-sidebar',
// init: async () => {
// try {
// const res = await api.fetchApi('/mtb/server-info')
// const msg = await res.json()
// exposed = msg.exposed
// } catch (e) {
// console.error('Error:', e)
// }
// },
init: () => {
let handle
const version = window?.__COMFYUI_FRONTEND_VERSION__
console.log(`%c ${version}`, 'background: orange; color: white;')
ensureMTBStyles()
app.ui.settings.addSetting({
id: 'mtb.io-sidebar.count',
category: ['mtb', 'Input & Output Sidebar', 'count'],
name: 'Number of images to fetch',
type: 'number',
defaultValue: 1000,
tooltip:
"This setting affects the input/output sidebar to determine how many images to fetch per pagination (pagination is not yet supported so for now it's the static total)",
attrs: {
style: {
// fontFamily: 'monospace',
},
},
})
app.ui.settings.addSetting({
id: 'mtb.io-sidebar.img-size',
category: ['mtb', 'Input & Output Sidebar', 'img-size'],
name: 'Resolution of the images',
type: 'number',
defaultValue: 512,
tooltip: "It's recommended to keep it at 512px",
attrs: {
style: {
// fontFamily: 'monospace',
},
},
})
app.ui.settings.addSetting({
id: 'mtb.io-sidebar.sort',
category: ['mtb', 'Input & Output Sidebar', 'sort'],
name: 'Default sort mode',
type: 'combo',
onChange: (v) => {
// alert(`Sort is now ${v}`)
currentSort = v
},
defaultValue: 'Modified',
// tooltip: "It's recommended to keep it at 512px",
options: [
'None',
'Modified',
'Modified-Reverse',
'Name',
'Name-Reverse',
],
})
app.extensionManager.registerSidebarTab({
id: 'mtb-inputs-outputs',
icon: 'pi pi-images',
title: 'Input & Outputs',
tooltip: 'MTB: Browse inputs and outputs directories.',
type: 'custom',
// this is run everytime the tab's diplay is toggled on.
render: async (el) => {
if (handle) {
handle.unregister()
handle = undefined
}
if (el.parentNode) {
el.parentNode.style.overflowY = 'clip'
}
const allModes = await getModes()
const input_modes = allModes.input.map((m) => `input - ${m}`)
const output_modes = allModes.output.map((m) => `output - ${m}`)
const urls = await getUrls()
let imgs = {}
const cont = makeElement('div.mtb_sidebar')
const imgGrid = makeElement('div.mtb_img_grid')
const selector = makeSelect(
['input', 'output', 'video', ...output_modes, ...input_modes],
currentMode,
)
selector.addEventListener('change', async (e) => {
let newMode = e.target.value
let changed = false
let newSub = ''
if (newMode !== 'input' && newMode !== 'output') {
if (newMode.startsWith('input - ')) {
newSub = newMode.replace('input - ', '')
newMode = 'input'
} else if (newMode.startsWith('output - ')) {
newSub = newMode.replace('output - ', '')
newMode = 'output'
}
}
changed = newMode !== currentMode || newSub !== subfolder
currentMode = newMode
subfolder = newSub
if (changed) {
imgGrid.innerHTML = ''
const urls = await getUrls(subfolder)
if (urls) {
imgs = getImgsFromUrls(urls, imgGrid)
}
}
})
const imgTools = makeElement('div.mtb_tools')
const orderSelect = makeSelect(
['None', 'Modified', 'Modified-Reverse', 'Name', 'Name-Reverse'],
currentSort,
)
orderSelect.addEventListener('change', async (e) => {
const newSort = e.target.value
const changed = newSort !== currentSort
currentSort = newSort
if (changed) {
imgGrid.innerHTML = ''
const urls = await getUrls(subfolder)
if (urls) {
imgs = getImgsFromUrls(urls, imgGrid)
}
}
})
const sizeSlider = makeSlider(64, 1024, currentWidth, 1)
imgTools.appendChild(orderSelect)
imgTools.appendChild(sizeSlider)
imgs = getImgsFromUrls(urls, imgGrid)
sizeSlider.addEventListener('input', (e) => {
currentWidth = e.target.value
for (const img of imgs) {
img.style.width = `${e.target.value}px`
}
})
handle = renderSidebar(el, cont, [selector, imgGrid, imgTools])
},
destroy: () => {
if (handle) {
handle.unregister()
handle = undefined
}
},
})
},
}
app.registerExtension(sidebar_extension)
}
+28
View File
@@ -0,0 +1,28 @@
// NOTE: this will be the LT part of mtb API system
// I need to properly publish the source and fix a few things before
// import { app } from '../../scripts/app.js'
// // import { api } from '../../scripts/api.js'
//
// import * as shared from './comfy_shared.js'
// import { createOutliner } from './dist/mtb_inspector.js'
//
// if (window?.__COMFYUI_FRONTEND_VERSION__) {
// const version = window?.__COMFYUI_FRONTEND_VERSION__
// console.log(`%c ${version}`, 'background: orange; color: white;')
//
// const panel = app.extensionManager.registerSidebarTab({
// id: 'mtb-nodes',
// icon: 'pi pi-bolt',
// title: 'MTB',
// tooltip: 'MTB: API outliner',
// type: 'custom',
// // this is run everytime the tab's diplay is toggled on.
// render: (el) => {
// const outliner = createOutliner(el)
// const inputs = shared.getAPIInputs()
// console.log('INPUTS', inputs)
// outliner.$$set({ inputs })
// },
// })
// }
+504
View File
@@ -0,0 +1,504 @@
/**
* Adds a named stylesheet to the document with an optional ability to replace an existing one.
*
* @param {string} name - The unique name (ID) of the stylesheet.
* @param {string} css - The CSS rules as a string.
* @param {boolean} [force=false] - Whether to replace the existing stylesheet if it exists.
* @returns {void}
*/
export function addNamedStyleSheet(name, css, force = false) {
const existingStyleSheet = document.getElementById(name)
if (existingStyleSheet && !force) {
console.debug(
`Stylesheet with name "${name}" already exists. Skipping addition.`,
)
return
}
if (existingStyleSheet && force) {
console.debug(`Stylesheet with name "${name}" exists. Replacing...`)
existingStyleSheet.remove()
}
const styleElement = document.createElement('style')
styleElement.id = name
styleElement.type = 'text/css'
styleElement.appendChild(document.createTextNode(css))
document.head.appendChild(styleElement)
console.debug(`Stylesheet with name "${name}" added.`)
}
export const ensureMTBStyles = () => {
const S = {
fg: 'var(--fg-color)',
bgi: 'var(--comfy-input-bg)',
bgm: 'var(--comfy-menu-bg)',
border: 'var(--comfy-border)',
borderHover: 'var(--comfy-border-hover)',
box: 'var(--comfy-box)',
accent: 'var(--p-button-text-primary-color)',
}
const common = `
.mtb_sidebar {
display: flex;
flex-direction: column;
background: ${S.bgm};
}
.mtb_img_grid {
display: flex;
flex-wrap: wrap;
overflow: scroll;
gap: 1em;
align-items: center;
justify-content: center;
height: 100%;
width: 100%;
}
.mtb_tools {
display: flex;
flex-direction: row;
align-items: center;
justify-content: space-between;
width: 100%;
}
`
const inputs = `
/* SELECT */
.mtb_select {
appearance: none;
display: grid;
grid-template-areas: "select";
padding: 10px;
background-color: ${S.bgi};
border: none;
border-radius: 5px;
font-size: 14px;
color: ${S.fg};
cursor: pointer;
width: 100%;
}
@supports (-moz-appearance:none) {
.mtb_select{
grid-area: select;
background: ${S.bgi} url('data:image/gif;base64,R0lGODlhBgAGAKEDAFVVVX9/f9TU1CgmNyH5BAEKAAMALAAAAAAGAAYAAAIODA4hCDKWxlhNvmCnGwUAOw==') right center no-repeat !important;
background-position: calc(100% - 5px) center !important;
-moz-appearance:none !important;
}
/* styling the dropdown arrow for browsers that support it */
.mtb_select:after {
content: "";
width: 0.8em;
height: 0.5em;
background-color: ${S.fg};
clip-path: polygon(100% 0%, 0 0%, 50% 100%);
}
.mtb_select:focus {
outline: none;
border-color: #0056b3;
}
.mtb_select > option {
padding: 10px;
background-color: ${S.bgi};
border:none;
color: ${S.fg};
}
.mtb_select > option:hover {
background-color: red;
color: ${S.fg};
}
/* SLIDER */
.mtb_slider[type="range"] {
-webkit-appearance: none;
appearance: none;
width: 100%;
height: 10px;
background: ${S.bgm};
border-radius: 5px;
outline: none;
opacity: 0.7;
transition: opacity .2s;
padding: 1em;
}
/* slider track */
.mtb_slider[type="range"]::-webkit-slider-runnable-track,
.mtb_slider[type="range"]::-moz-range-track {
width: 100%;
height: 10px;
background: ${S.bgi};
border-radius: 5px;
}
/* progress */
.mtb_slider[type="range"]::-moz-range-progress {
background-color: ${S.accent};
height:10px;
border-radius: 5px;
}
/* slider thumb (the handle) */
.mtb_slider[type="range"]::-webkit-slider-thumb,
.mtb_slider[type="range"]::-moz-range-thumb
{
-webkit-appearance: none;
appearance: none;
width: 15px;
height: 15px;
border-radius: 50%;
background: ${S.fg};
border: none;
cursor: pointer;
filter: drop-shadow(1px 1px 4px black);
}
.mtb_slider[type="range"]:focus {
opacity: 1;
}
.mtb_slider[type=range]:-moz-focusring{
outline: 1px solid red;
outline-offset: -1px;
}
.mtb_slider[type="range"]:hover::-webkit-slider-thumb,
.mtb_slider[type="range"]:active::-webkit-slider-thumb {
background-color: ${S.accent};
}
`
addNamedStyleSheet(
'mtb_ui',
`
${common}
${inputs}
`,
)
}
/**
* Creates a DOM element with optional styles, class, and id.
*
* @param {string} kind - The tag name of the element. Supports class and id syntax (e.g. 'div.class#id').
* @param {Object} [style] - CSS styles to apply to the element.
* @returns {HTMLElement} - The created DOM element.
*/
export const makeElement = (kind, style) => {
let [real_kind, className] = kind.split('.')
let id
if (className?.includes('#')) {
;[className, id] = className.split('#')
}
const el = document.createElement(real_kind)
if (style) {
Object.assign(el.style, style)
}
if (className) {
el.classList.add(...className.split(' ')) // Support multiple classes
}
if (id) {
el.id = id
}
return el
}
/**
* Clears all child elements of the given parent element.
*
* @param {HTMLElement} el - The parent element whose children should be removed.
*/
export const clearElement = (el) => {
while (el.firstChild) {
el.removeChild(el.firstChild)
}
}
/**
* Creates a labeled element (input, select, etc.).
*
* @param {HTMLElement} el - The element to label.
* @param {string} labelText - The label text.
* @returns {HTMLDivElement} - A div containing the label and the element.
*/
export const makeLabeledElement = (el, labelText) => {
const wrapper = makeElement('div.mtb_labeled_element', {
marginBottom: '1em',
})
const label = makeElement('label', {
display: 'block',
marginBottom: '0.5em',
})
label.textContent = labelText
wrapper.appendChild(label)
wrapper.appendChild(el)
return wrapper
}
/**
* Converts a camelCase CSS property to kebab-case.
*
* @param {string} prop - The camelCase CSS property.
* @returns {string} - The kebab-case CSS property.
*/
const camelToKebab = (prop) =>
prop.replace(/[A-Z]/g, (match) => `-${match.toLowerCase()}`)
/**
* Parses the style string into an object of CSS property-value pairs.
*
* @param {string} styleString - The CSS rule text (e.g., "color: red; background-color: blue;").
* @returns {Object} - An object with camelCase CSS properties.
*/
const parseStyleString = (styleString) => {
const styleObj = {}
for (const rule of styleString.split(';')) {
const [property, value] = rule.split(':').map((item) => item.trim())
if (property && value) {
const camelProp = property.replace(/-([a-z])/g, (g) => g[1].toUpperCase())
styleObj[camelProp] = value
}
}
return styleObj
}
/**
* Defines a new CSS class with the provided styles, or skips if the class already exists.
*
* @param {string} className - The name of the CSS class to define.
* @param {Object} classStyles - An object containing camelCase CSS property-value pairs.
*/
export function defineCSSClass(className, classStyles) {
const styleSheets = document.styleSheets
let classExists = false
let existingStyleString = ''
const classExistsInStyleSheet = (styleSheet) => {
const rules = styleSheet.rules || styleSheet.cssRules
for (const rule of rules) {
if (rule.selectorText === `.${className}`) {
classExists = true
existingStyleString = rule.style.cssText // Capture existing styles
return true
}
}
return false
}
for (const styleSheet of styleSheets) {
if (classExistsInStyleSheet(styleSheet)) {
console.debug(`Class ${className} already exists, merging styles...`)
break
}
}
const existingStyles = classExists
? parseStyleString(existingStyleString)
: {}
const mergedStyles = { ...existingStyles, ...classStyles }
const stylesString = Object.entries(mergedStyles)
.map(([key, value]) => `${camelToKebab(key)}: ${value};`)
.join(' ')
if (!classExists) {
console.debug(`Defining new class ${className}...`)
if (styleSheets[0].insertRule) {
styleSheets[0].insertRule(`.${className} { ${stylesString} }`, 0)
} else if (styleSheets[0].addRule) {
styleSheets[0].addRule(`.${className}`, stylesString, 0)
}
} else {
console.debug(`Updating existing class ${className} with merged styles...`)
for (const styleSheet of styleSheets) {
const rules = styleSheet.rules || styleSheet.cssRules
for (const rule of rules) {
if (rule.selectorText === `.${className}`) {
rule.style.cssText = stylesString // Update the existing rule
}
}
}
}
console.debug(
`Class ${className} has been defined/updated with styles:`,
mergedStyles,
)
}
/**
* Renders a sidebar and ensures it resizes correctly when the window is resized.
*
* @param {HTMLElement} el - The element where the sidebar is rendered.
* @param {HTMLElement} cont - The content container of the sidebar.
* @param {HTMLElement[]} elems - Array of elements to append to the sidebar.
* @returns {Object} - A handle with a method to unregister the resize event.
*/
export const renderSidebar = (el, cont, elems) => {
el.appendChild(cont)
if (!el.parentNode) {
return
}
el.parentNode.style.overflowY = 'clip'
cont.style.height = `${el.parentNode.offsetHeight}px`
const resizeHandler = () => {
cont.style.height = `${el.parentNode.offsetHeight}px`
}
window.addEventListener('resize', resizeHandler)
for (const elem of elems) {
cont.appendChild(elem)
}
return {
unregister: () => {
window.removeEventListener('resize', resizeHandler)
},
}
}
/**
* Creates a <select> dropdown with given options.
*
* @param {string[]} options - The options for the select element.
* @param {string} [current] - The currently selected option (optional).
* @returns {HTMLSelectElement} - The created <select> element.
*/
export const makeSelect = (options, current = undefined) => {
const selector = makeElement('select.mtb_select', {
width: 'auto',
margin: '1em',
})
for (const option of options) {
const opt = makeElement('option')
opt.value = option
opt.innerHTML = option
selector.appendChild(opt)
}
if (current !== undefined) {
if (options.includes(current)) {
selector.value = current
} else {
console.error(
`You tried to select an option that doesn't exist (${current}). Options: ${options}`,
)
}
}
return selector
}
/**
* Creates an <input type="range"> slider element with given parameters.
*
* @param {number} min - Minimum value of the slider.
* @param {number} max - Maximum value of the slider.
* @param {number} [value] - Initial value of the slider.
* @param {number} [step] - Step value for the slider.
* @returns {HTMLInputElement} - The created slider element.
*/
export const makeSlider = (min, max, value = undefined, step = undefined) => {
const slider = makeElement('input.mtb_slider', {
width: '100%',
})
slider.type = 'range'
slider.min = min || 0
slider.max = max || 100
slider.value = value || slider.min
slider.step = step || 1
return slider
}
/**
* Creates a button element.
*
* @param {string} label - The label for the button.
* @param {Object} [style] - Optional styles to apply to the button.
* @param {Function} [onClick] - Optional click handler.
* @returns {HTMLButtonElement} - The created button element.
*/
export const makeButton = (label, style = {}, onClick = undefined) => {
const button = makeElement('button.mtb_button', style)
button.textContent = label
if (onClick) {
button.addEventListener('click', onClick)
}
return button
}
/**
* Creates a resizable splitter between two elements.
*
* @param {HTMLElement} el1 - The first element.
* @param {HTMLElement} el2 - The second element.
* @param {'vertical' | 'horizontal'} direction - Splitter direction (vertical or horizontal).
* @param {'absolute' | 'normal'} mode - Splitter mode: 'absolute' for free resizing, 'normal' for layout-based resizing.
* @returns {HTMLDivElement} - The container with resizable splitter.
*/
export const makeSplitter = (
el1,
el2,
direction = 'vertical',
mode = 'normal',
) => {
const container = makeElement('div.mtb_splitter_container', {
display: mode === 'absolute' ? 'block' : 'flex',
flexDirection: direction === 'vertical' ? 'row' : 'column',
position: mode === 'absolute' ? 'relative' : 'static',
height: '100%',
width: '100%',
})
const handle = makeElement('div.mtb_splitter_handle', {
backgroundColor: '#ccc',
cursor: direction === 'vertical' ? 'col-resize' : 'row-resize',
width: direction === 'vertical' ? '5px' : '100%',
height: direction === 'horizontal' ? '5px' : '100%',
})
let isResizing = false
handle.addEventListener('mousedown', () => {
isResizing = true
})
window.addEventListener('mouseup', () => {
isResizing = false
})
window.addEventListener('mousemove', (e) => {
if (!isResizing) return
if (direction === 'vertical') {
const newWidth = e.clientX - container.offsetLeft
el1.style.width = `${newWidth}px`
el2.style.width = `${container.offsetWidth - newWidth}px`
} else {
const newHeight = e.clientY - container.offsetTop
el1.style.height = `${newHeight}px`
el2.style.height = `${container.offsetHeight - newHeight}px`
}
})
container.appendChild(el1)
container.appendChild(handle)
container.appendChild(el2)
return container
}
+152 -92
View File
@@ -14,13 +14,14 @@
import { app } from '../../scripts/app.js' import { app } from '../../scripts/app.js'
import { api } from '../../scripts/api.js' import { api } from '../../scripts/api.js'
import * as mtb_ui from './mtb_ui.js'
import parseCss from './extern/parse-css.js' import parseCss from './extern/parse-css.js'
import * as shared from './comfy_shared.js' import * as shared from './comfy_shared.js'
import { infoLogger } from './comfy_shared.js' import { infoLogger } from './comfy_shared.js'
import { NumberInputWidget } from './numberInput.js' import { NumberInputWidget } from './numberInput.js'
// NOTE: new widget types registered by MTB Widgets // NOTE: new widget types registered by MTB Widgets
const newTypes = [, /*'BOOL'*/ 'COLOR', 'BBOX'] const newTypes = [/*'BOOL'*/ , 'COLOR', 'BBOX']
const deprecated_nodes = { const deprecated_nodes = {
// 'Animation Builder': // 'Animation Builder':
@@ -96,7 +97,7 @@ export function addVectorWidgetW(
name, name,
value, value,
vector_size, vector_size,
callback, _callback,
app, app,
) { ) {
// const inputEl = document.createElement('div') // const inputEl = document.createElement('div')
@@ -243,7 +244,7 @@ export const MtbWidgets = {
y: 0, y: 0,
options: { default: Array.from({ length: size }, () => 0.0) }, options: { default: Array.from({ length: size }, () => 0.0) },
_value: val || Array.from({ length: size }, () => 0.0), _value: val || Array.from({ length: size }, () => 0.0),
draw: function (ctx, node, width, widgetY, height) { draw: (ctx, node, width, widgetY, height) => {
ctx.textAlign = 'left' ctx.textAlign = 'left'
ctx.strokeStyle = outline_color ctx.strokeStyle = outline_color
ctx.fillStyle = background_color ctx.fillStyle = background_color
@@ -311,7 +312,7 @@ export const MtbWidgets = {
value: val?.default || [0, 0, 0, 0], value: val?.default || [0, 0, 0, 0],
options: {}, options: {},
draw: function (ctx, node, widget_width, widgetY, height) { draw: function (ctx, _node, widget_width, widgetY, _height) {
const hide = this.type !== 'BBOX' && app.canvas.ds.scale > 0.5 const hide = this.type !== 'BBOX' && app.canvas.ds.scale > 0.5
const show_text = true const show_text = true
@@ -321,13 +322,13 @@ export const MtbWidgets = {
const secondary_text_color = LiteGraph.WIDGET_SECONDARY_TEXT_COLOR const secondary_text_color = LiteGraph.WIDGET_SECONDARY_TEXT_COLOR
const H = LiteGraph.NODE_WIDGET_HEIGHT const H = LiteGraph.NODE_WIDGET_HEIGHT
let margin = 15 const margin = 15
let numWidgets = 4 // Number of stacked widgets const numWidgets = 4 // Number of stacked widgets
if (hide) return if (hide) return
for (let i = 0; i < numWidgets; i++) { for (let i = 0; i < numWidgets; i++) {
let currentY = widgetY + i * (H + margin) // Adjust Y position for each widget const currentY = widgetY + i * (H + margin) // Adjust Y position for each widget
ctx.textAlign = 'left' ctx.textAlign = 'left'
ctx.strokeStyle = outline_color ctx.strokeStyle = outline_color
@@ -535,21 +536,34 @@ export const MtbWidgets = {
picker.type = 'color' picker.type = 'color'
picker.value = this.value picker.value = this.value
picker.style.position = 'absolute' Object.assign(picker.style, {
picker.style.left = '999999px' //(window.innerWidth / 2) + "px"; position: 'fixed',
picker.style.top = '999999px' //(window.innerHeight / 2) + "px"; left: `${e.clientX}px`,
top: `${e.clientY}px`,
height: '0px',
width: '0px',
padding: '0px',
opacity: 0,
})
picker.addEventListener('blur', () => {
this.callback?.(this.value)
node.graph._version++
picker.remove()
})
picker.addEventListener('input', () => {
if (!picker.value) return
this.value = picker.value
app.canvas.setDirty(true)
})
document.body.appendChild(picker) document.body.appendChild(picker)
picker.addEventListener('change', () => { requestAnimationFrame(() => {
this.value = picker.value picker.showPicker()
this.callback?.(this.value) picker.focus()
node.graph._version++
node.setDirtyCanvas(true, true)
picker.remove()
}) })
picker.click()
} }
} }
} }
@@ -658,12 +672,11 @@ const mtb_widgets = {
init: async () => { init: async () => {
infoLogger('Registering mtb.widgets') infoLogger('Registering mtb.widgets')
try { try {
const res = await api.fetchApi('/mtb/debug') const msg = await shared.getServerInfo()
const msg = await res.json()
if (!window.MTB) { if (!window.MTB) {
window.MTB = {} window.MTB = {}
} }
window.MTB.DEBUG = msg.enabled window.MTB.DEBUG = msg.debug
} catch (e) { } catch (e) {
console.error('Error:', e) console.error('Error:', e)
} }
@@ -671,16 +684,26 @@ const mtb_widgets = {
setup: () => { setup: () => {
app.ui.settings.addSetting({ app.ui.settings.addSetting({
id: 'mtb.Debug.enabled', id: 'mtb.postshot.path',
name: '[⚡mtb] Enable Debug (py and js)', category: ['mtb', 'PostShot', 'path'],
name: 'Path to Postshot CLI',
type: 'string',
defaultValue: 'C:/Program Files/Jawset Postshot/bin/postshot-cli.exe',
tooltip: 'The path to the postshot CLI',
})
app.ui.settings.addSetting({
id: 'mtb.Main.debug-enabled',
category: ['mtb', 'Main', 'debug-enabled'],
name: 'Enable Debug (py and js)',
type: 'boolean', type: 'boolean',
defaultValue: false, defaultValue: false,
tooltip: tooltip:
'This will enable debug messages in the console and in the python console respectively', 'This will enable debug messages in the console and in the python console respectively, no need to restart the server, but do reload the webui',
attrs: { attrs: {
style: { style: {
fontFamily: 'monospace', // fontFamily: 'monospace',
}, },
}, },
async onChange(value) { async onChange(value) {
@@ -692,17 +715,11 @@ const mtb_widgets = {
infoLogger('Enabled DEBUG mode') infoLogger('Enabled DEBUG mode')
} }
await api try {
.fetchApi('/mtb/debug', { shared.setServerInfo({ debug: value })
method: 'POST', } catch (err) {
body: JSON.stringify({ console.error('Error:', err)
enabled: value, }
}),
})
.then((_response) => {})
.catch((error) => {
console.error('Error:', error)
})
}, },
}) })
}, },
@@ -751,7 +768,7 @@ const mtb_widgets = {
// const rinputs = nodeData.input?.required // const rinputs = nodeData.input?.required
let has_custom = false let has_custom = false
if (nodeData.input && nodeData.input.required) { if (nodeData.input?.required) {
for (const i of Object.keys(nodeData.input.required)) { for (const i of Object.keys(nodeData.input.required)) {
const input_type = nodeData.input.required[i][0] const input_type = nodeData.input.required[i][0]
@@ -764,10 +781,8 @@ const mtb_widgets = {
if (has_custom) { if (has_custom) {
//- Add widgets on node creation //- Add widgets on node creation
const onNodeCreated = nodeType.prototype.onNodeCreated const onNodeCreated = nodeType.prototype.onNodeCreated
nodeType.prototype.onNodeCreated = function () { nodeType.prototype.onNodeCreated = function (...args) {
const r = onNodeCreated const r = onNodeCreated ? onNodeCreated.apply(this, args) : undefined
? onNodeCreated.apply(this, arguments)
: undefined
this.serialize_widgets = true this.serialize_widgets = true
this.setSize?.(this.computeSize()) this.setSize?.(this.computeSize())
@@ -785,8 +800,8 @@ const mtb_widgets = {
? origGetExtraMenuOptions.apply(this, arguments) ? origGetExtraMenuOptions.apply(this, arguments)
: undefined : undefined
if (this.widgets) { if (this.widgets) {
let toInput = [] const toInput = []
let toWidget = [] const toWidget = []
for (const w of this.widgets) { for (const w of this.widgets) {
if (w.type === shared.CONVERTED_TYPE) { if (w.type === shared.CONVERTED_TYPE) {
//- This is already handled by widgetinputs.js //- This is already handled by widgetinputs.js
@@ -856,6 +871,22 @@ const mtb_widgets = {
break break
} }
case 'Postshot Train (mtb)':
case 'Postshot Export (mtb)': {
const onNodeCreated = nodeType.prototype.onNodeCreated
nodeType.prototype.onNodeCreated = function (...args) {
const r = onNodeCreated ? onNodeCreated.apply(this, args) : undefined
const { postshot_cli } = shared.getNamedWidget(this, 'postshot_cli')
shared.hideWidgetForGood(this, postshot_cli)
api.getSetting('mtb.postshot.path').then((p) => {
postshot_cli._value = p
})
}
break
}
case 'Save Gif (mtb)': case 'Save Gif (mtb)':
case 'Save Animated Image (mtb)': { case 'Save Animated Image (mtb)': {
const onExecuted = nodeType.prototype.onExecuted const onExecuted = nodeType.prototype.onExecuted
@@ -878,7 +909,7 @@ const mtb_widgets = {
imgURLs = imgURLs.concat( imgURLs = imgURLs.concat(
message.gif.map((params) => { message.gif.map((params) => {
return api.apiURL( return api.apiURL(
'/view?' + new URLSearchParams(params).toString(), `/view?${new URLSearchParams(params).toString()}`,
) )
}), }),
) )
@@ -887,7 +918,7 @@ const mtb_widgets = {
imgURLs = imgURLs.concat( imgURLs = imgURLs.concat(
message.apng.map((params) => { message.apng.map((params) => {
return api.apiURL( return api.apiURL(
'/view?' + new URLSearchParams(params).toString(), `/view?${new URLSearchParams(params).toString()}`,
) )
}), }),
) )
@@ -915,37 +946,71 @@ const mtb_widgets = {
} }
case 'Animation Builder (mtb)': { case 'Animation Builder (mtb)': {
const onNodeCreated = nodeType.prototype.onNodeCreated const onNodeCreated = nodeType.prototype.onNodeCreated
nodeType.prototype.onNodeCreated = function () { nodeType.prototype.onNodeCreated = function (...args) {
const r = onNodeCreated const r = onNodeCreated ? onNodeCreated.apply(this, args) : undefined
? onNodeCreated.apply(this, arguments)
: undefined
this.changeMode(LiteGraph.ALWAYS) this.changeMode(LiteGraph.ALWAYS)
const { raw_iteration, raw_loop, total_frames, loop_count } =
const raw_iteration = this.widgets.find( shared.getNamedWidget(
(w) => w.name === 'raw_iteration', this,
) 'raw_iteration',
const raw_loop = this.widgets.find((w) => w.name === 'raw_loop') 'raw_loop',
'total_frames',
const total_frames = this.widgets.find( 'loop_count',
(w) => w.name === 'total_frames', )
)
const loop_count = this.widgets.find((w) => w.name === 'loop_count')
shared.hideWidgetForGood(this, raw_iteration) shared.hideWidgetForGood(this, raw_iteration)
shared.hideWidgetForGood(this, raw_loop) shared.hideWidgetForGood(this, raw_loop)
raw_iteration._value = 0 raw_iteration._value = 0
const value_preview = this.addCustomWidget( // const value_preview = this.addCustomWidget(
MtbWidgets['DEBUG_STRING']('value_preview', 'Idle'), // MtbWidgets.DEBUG_STRING('value_preview', 'Idle'),
) // )
value_preview.parent = this
const loop_preview = this.addCustomWidget( const dom_value_preview = mtb_ui.makeElement('p', {
MtbWidgets['DEBUG_STRING']('loop_preview', 'Iteration: Idle'), fontWeigth: '700',
textAlign: 'center',
fontSize: '1.5em',
margin: 0,
})
const value_preview = this.addDOMWidget(
'value_preview',
'DISPLAY',
dom_value_preview,
{
hideOnZoom: false,
setValue: (val) => {
if (val) {
value_preview.element.innerHTML = val
}
},
},
) )
loop_preview.parent = this value_preview.value = 'Idle'
const dom_loop_preview = mtb_ui.makeElement('p', {
textAlign: 'center',
margin: 0,
})
const loop_preview = this.addDOMWidget(
'loop_preview',
'DISPLAY',
dom_loop_preview,
{
hideOnZoom: false,
setValue: (val) => {
if (val) {
dom_loop_preview.innerHTML = val
}
},
getValue: () => {
dom_loop_preview.innerHTML
},
},
)
loop_preview.value = 'Iteration: Idle'
const onReset = () => { const onReset = () => {
raw_iteration.value = 0 raw_iteration.value = 0
@@ -958,10 +1023,10 @@ const mtb_widgets = {
} }
// reset button // reset button
this.addWidget('button', `Reset`, 'reset', onReset) this.addWidget('button', 'Reset', 'reset', onReset)
// run button // run button
this.addWidget('button', `Queue`, 'queue', () => { this.addWidget('button', 'Queue', 'queue', () => {
onReset() // this could maybe be a setting or checkbox onReset() // this could maybe be a setting or checkbox
app.queuePrompt(0, total_frames.value * loop_count.value) app.queuePrompt(0, total_frames.value * loop_count.value)
window.MTB?.notify?.( window.MTB?.notify?.(
@@ -1001,9 +1066,9 @@ const mtb_widgets = {
} }
case 'Interpolate Clip Sequential (mtb)': { case 'Interpolate Clip Sequential (mtb)': {
const onNodeCreated = nodeType.prototype.onNodeCreated const onNodeCreated = nodeType.prototype.onNodeCreated
nodeType.prototype.onNodeCreated = function () { nodeType.prototype.onNodeCreated = function (...args) {
const r = onNodeCreated const r = onNodeCreated
? onNodeCreated.apply(this, arguments) ? onNodeCreated.apply(this, ...args)
: undefined : undefined
const addReplacement = () => { const addReplacement = () => {
const input = this.addInput( const input = this.addInput(
@@ -1015,19 +1080,14 @@ const mtb_widgets = {
this.addWidget('STRING', `replacement_${this.widgets.length}`, '') this.addWidget('STRING', `replacement_${this.widgets.length}`, '')
} }
//- add //- add
this.addWidget('button', '+', 'add', function (value, widget, node) { this.addWidget('button', '+', 'add', (value, widget, node) => {
console.log('Button clicked', value, widget, node) console.log('Button clicked', value, widget, node)
addReplacement() addReplacement()
}) })
//- remove //- remove
this.addWidget( this.addWidget('button', '-', 'remove', (value, widget, node) => {
'button', console.log(`Button clicked: ${value}`, widget, node)
'-', })
'remove',
function (value, widget, node) {
console.log(`Button clicked: ${value}`, widget, node)
},
)
return r return r
} }
@@ -1042,16 +1102,10 @@ const mtb_widgets = {
const getStyle = async (node) => { const getStyle = async (node) => {
try { try {
const getStyles = await api.fetchApi('/mtb/actions', { const getStyles = await runAction(
method: 'POST', 'getStyles',
body: JSON.stringify({ node.widgets?.[0].value ? node.widgets[0].value : '',
name: 'getStyles', )
args:
node.widgets && node.widgets[0].value
? node.widgets[0].value
: '',
}),
})
const output = await getStyles.json() const output = await getStyles.json()
return output?.result return output?.result
@@ -1123,6 +1177,10 @@ const mtb_widgets = {
shared.setupDynamicConnections(nodeType, 'video', 'VIDEO') shared.setupDynamicConnections(nodeType, 'video', 'VIDEO')
break break
} }
case 'Interpolate Condition (mtb)': {
shared.setupDynamicConnections(nodeType, 'condition', 'CONDITIONING')
break
}
case 'Psd Save (mtb)': { case 'Psd Save (mtb)': {
shared.setupDynamicConnections(nodeType, 'input_', 'PSDLAYER') shared.setupDynamicConnections(nodeType, 'input_', 'PSDLAYER')
break break
@@ -1147,6 +1205,8 @@ const mtb_widgets = {
shared.setupDynamicConnections(nodeType, 'floats', 'FLOATS') shared.setupDynamicConnections(nodeType, 'floats', 'FLOATS')
break break
} }
case 'Batch Sequence (mtb)':
case 'Batch Sequence Plus (mtb)':
case 'Batch Merge (mtb)': { case 'Batch Merge (mtb)': {
shared.setupDynamicConnections(nodeType, 'batches', 'IMAGE') shared.setupDynamicConnections(nodeType, 'batches', 'IMAGE')
@@ -1159,13 +1219,13 @@ const mtb_widgets = {
const r = onNodeCreated const r = onNodeCreated
? onNodeCreated.apply(this, arguments) ? onNodeCreated.apply(this, arguments)
: undefined : undefined
this.addInput(`x`, '*') this.addInput('x', '*')
return r return r
} }
const onConnectionsChange = nodeType.prototype.onConnectionsChange const onConnectionsChange = nodeType.prototype.onConnectionsChange
nodeType.prototype.onConnectionsChange = function ( nodeType.prototype.onConnectionsChange = function (
type, _type,
index, index,
connected, connected,
link_info, link_info,
@@ -1180,7 +1240,7 @@ const mtb_widgets = {
//- infer type //- infer type
if (link_info) { if (link_info) {
const fromNode = this.graph._nodes.find( const fromNode = this.graph._nodes.find(
(otherNode) => otherNode.id == link_info.origin_id, (otherNode) => otherNode.id !== link_info.origin_id,
) )
const type = fromNode.outputs[link_info.origin_slot].type const type = fromNode.outputs[link_info.origin_slot].type
this.inputs[index].type = type this.inputs[index].type = type
+246
View File
@@ -0,0 +1,246 @@
// web/note_plus.constants.js
export const DEFAULT_CSS = ''
export const DEFAULT_HTML = `<p style='color:red;font-family:monospace'>
Note+
</p>`
export const DEFAULT_MD = '## Note+'
export const DEFAULT_MODE = 'markdown'
export const DEFAULT_THEME = 'one_dark'
export const DEMO_CONTENT = `
# @mtb/svelte-markdown.
## This is a subheader
[![embedded test](https://github.com/melMass/comfy_mtb/actions/workflows/test_embedded.yml/badge.svg)](https://github.com/melMass/comfy_mtb/actions/workflows/test_embedded.yml)
![home](https://repository-images.githubusercontent.com/649047066/a3eef9a7-20dd-4ef9-b839-884502d4e873)
<details>
<summary>More details about the inception of the project</summary>
\`\`\`js
class YesMan{
constructor(){
this.started = false
}
}
\`\`\`
</details>
This is a paragraph. If it goes over the maximum width it will not automatically wrap unless it reaches the max-w of \`prose\` check [styles](/styles) for more info.
This component is useful for building some tools on top. Or even just a static system using svelte at its core. My personal blog is fully powered by **@mtb/svelte-markdown**
| And this is | A table |
|-------------|---------|
| With two | columns |
We also support github callout:
> [!NOTE]
> Highlights information that users should take into account, even when skimming.
> [!TIP]
> Optional information to help a user be more successful.
> [!IMPORTANT]
> Crucial information necessary for users to succeed.
> [!WARNING]
> Critical content demanding immediate user attention due to potential risks.
> [!CAUTION]
> Negative potential consequences of an action.
`
export const THEMES = [
'ambiance',
'chaos',
'chrome',
'cloud9_day',
'cloud9_night',
'cloud9_night_low_color',
'cloud_editor',
'cloud_editor_dark',
'clouds',
'clouds_midnight',
'cobalt',
'crimson_editor',
'dawn',
'dracula',
'dreamweaver',
'eclipse',
'github',
'github_dark',
'gob',
'gruvbox',
'gruvbox_dark_hard',
'gruvbox_light_hard',
'idle_fingers',
'iplastic',
'katzenmilch',
'kr_theme',
'kuroir',
'merbivore',
'merbivore_soft',
'mono_industrial',
'monokai',
'nord_dark',
'one_dark',
'pastel_on_dark',
'solarized_dark',
'solarized_light',
'sqlserver',
'terminal',
'textmate',
'tomorrow',
'tomorrow_night',
'tomorrow_night_blue',
'tomorrow_night_bright',
'tomorrow_night_eighties',
'twilight',
'vibrant_ink',
'vscode',
]
export const CSS_RESET = `
* {
font-family: monospace;
line-height: 1.25em;
}
.shiki{
padding: 1em;
width: 100%;
}
.markdown-callout-title {
.octicon{
fill:white;
}
/* background: var(--current-color); */
color: var(--current-color);
font-weight: bold;
/* border-start-end-radius: var(--radius); */
/* border-start-start-radius: var(--radius); */
padding: 0.5em;
padding-inline-start: 1em;
}
.markdown-callout-content {
padding: 1em;
}
.markdown-callout {
--radius: 8px;
--current-color: purple;
/* border-start-end-radius: var(--radius); */
/* border-start-start-radius: var(--radius); */
border-left: 3px solid var(--current-color);
margin-bottom: 1em;
margin-top: 1em;
}
.markdown-callout-tip {
--text-color: whitesmoke;
--current-color: #50e3c2;
}
.markdown-callout-note {
--text-color: whitesmoke;
--current-color: #0070f3;
}
.markdown-callout-important {
--text-color: whitesmoke;
--current-color: #7928ca;
}
.markdown-callout-warning {
--current-color: #f5a623;
}
.markdown-callout-caution {
--current-color: #e60000;
}
.note-plus-preview {
display:flex;
flex-direction:column;
align-items: flex-start;
width:95%;
margin-left: 20px;
margin-top:20px;
/*background-color: rgba(255,0,0,0.5)!important;*/
}
/* allowed to be selected*/
h1, h2, h3, h4, h5, h6,a, p, ul, ol, dl, blockquote,details,summary {
pointer-events:auto;
user-select:text;
}
h1, h2, h3, h4, h5, h6 {
display:inline-block;
margin: 0;
padding: 0;
font-weight: normal;
}
p, ul, ol, dl, blockquote {
margin: 0.3em;
padding: 0;
}
ul, ol {
padding-left: 1em;
}
a {
color: inherit;
text-decoration: none;
pointer-events: all;
color: cyan;
}
img {
padding: 1em 0;
max-width: 100%;
}
iframe {
max-width: 100%;
height: auto;
border:none;
pointer-events:all;
}
blockquote {
border-left: 4px solid #ccc;
padding-left: 1em;
margin-left: 0;
font-style: italic;
}
pre, code {
font-family: monospace;
}
table {
border-collapse: collapse;
width: 100%;
border-bottom: 1px solid #000;
margin: 1em 0;
}
th, td {
border-left: 1px solid #000;
border-right: 1px solid #000;
padding: 8px;
text-align: left;
}
th {
border: 1px solid #000;
background-color: rgba(0,0,0,0.5);
}
input[type="checkbox"] {
margin-right: 10px;
}
`
+390 -253
View File
@@ -1,155 +1,122 @@
/// <reference path="../types/typedefs.js" />
import { app } from '../../scripts/app.js' import { app } from '../../scripts/app.js'
import * as shared from './comfy_shared.js' import * as shared from './comfy_shared.js'
import { infoLogger, successLogger, errorLogger } from './comfy_shared.js' import { infoLogger, successLogger, errorLogger } from './comfy_shared.js'
import {
DEFAULT_CSS,
DEFAULT_HTML,
DEFAULT_MD,
DEFAULT_MODE,
DEFAULT_THEME,
THEMES,
CSS_RESET,
DEMO_CONTENT,
} from './note_plus.constants.js'
import { LocalStorageManager } from './comfy_shared.js'
const DEFAULT_CSS = '' const storage = new LocalStorageManager('mtb')
const DEFAULT_HTML = `<p style='color:red;font-family:monospace'>
Note+
</p>`
const DEFAULT_MD = '## Note+'
const DEFAULT_MODE = 'markdown'
const DEFAULT_THEME = 'one_dark'
const CSS_RESET = ` /**
* { * Uses `@mtb/markdown-parser` (a fork of marked)
font-family: monospace; * It is statically stored to avoid having
line-height: 1.25em; * more than 1 instance ever.
* The size difference between both libraries...
* ╭───┬────────────────────────────────┬──────────╮
* │ # │ name │ size │
* ├───┼────────────────────────────────┼──────────┤
* │ 0 │ web-dist/mtb_markdown_plus.mjs │ 1.2 MB │ <- with shiki
* │ 1 │ web-dist/mtb_markdown.mjs │ 44.7 KB │
* ╰───┴────────────────────────────────┴──────────╯
*/
let useShiki = storage.get('np-use-shiki', false)
const makeResizable = (dialog) => {
dialog.style.resize = 'both'
dialog.style.transformOrigin = 'top left'
dialog.style.overflow = 'auto'
} }
h1, h2, h3, h4, h5, h6 { const makeDraggable = (dialog, handle) => {
margin: 0; let offsetX = 0
padding: 0; let offsetY = 0
font-weight: normal; let isDragging = false
const onMouseMove = (e) => {
if (isDragging) {
dialog.style.left = `${e.clientX - offsetX}px`
dialog.style.top = `${e.clientY - offsetY}px`
}
}
const onMouseUp = () => {
isDragging = false
document.removeEventListener('mousemove', onMouseMove)
document.removeEventListener('mouseup', onMouseUp)
}
handle.addEventListener('mousedown', (e) => {
isDragging = true
offsetX = e.clientX - dialog.offsetLeft
offsetY = e.clientY - dialog.offsetTop
document.addEventListener('mousemove', onMouseMove)
document.addEventListener('mouseup', onMouseUp)
})
} }
p, ul, ol, dl, blockquote { /** @extends {LGraphNode} */
margin: 0.3em;
padding: 0;
}
ul, ol {
padding-left: 1em;
}
a {
color: inherit;
text-decoration: none;
pointer-events: all;
color: cyan;
}
img {
padding: 1em 0;
max-width: 100%;
}
iframe {
width: 100%;
height: auto;
border:none;
pointer-events:all;
}
blockquote {
border-left: 4px solid #ccc;
padding-left: 1em;
margin-left: 0;
font-style: italic;
}
pre, code {
font-family: monospace;
}
table {
border-collapse: collapse;
width: 100%;
border-bottom: 1px solid #000;
margin: 1em 0;
}
th, td {
border-left: 1px solid #000;
border-right: 1px solid #000;
padding: 8px;
text-align: left;
}
th {
border: 1px solid #000;
background-color: rgba(0,0,0,0.5);
}
input[type="checkbox"] {
margin-right: 10px;
}
`
const themes = [
'ambiance',
'chaos',
'chrome',
'cloud9_day',
'cloud9_night',
'cloud9_night_low_color',
'cloud_editor',
'cloud_editor_dark',
'clouds',
'clouds_midnight',
'cobalt',
'crimson_editor',
'dawn',
'dracula',
'dreamweaver',
'eclipse',
'github',
'github_dark',
'gob',
'gruvbox',
'gruvbox_dark_hard',
'gruvbox_light_hard',
'idle_fingers',
'iplastic',
'katzenmilch',
'kr_theme',
'kuroir',
'merbivore',
'merbivore_soft',
'mono_industrial',
'monokai',
'nord_dark',
'one_dark',
'pastel_on_dark',
'solarized_dark',
'solarized_light',
'sqlserver',
'terminal',
'textmate',
'tomorrow',
'tomorrow_night',
'tomorrow_night_blue',
'tomorrow_night_bright',
'tomorrow_night_eighties',
'twilight',
'vibrant_ink',
'vscode',
]
class NotePlus extends LiteGraph.LGraphNode { class NotePlus extends LiteGraph.LGraphNode {
// same values as the comfy note // same values as the comfy note
color = LGraphCanvas.node_colors.yellow.color color = LGraphCanvas.node_colors.yellow.color
bgcolor = LGraphCanvas.node_colors.yellow.bgcolor bgcolor = LGraphCanvas.node_colors.yellow.bgcolor
groupcolor = LGraphCanvas.node_colors.yellow.groupcolor groupcolor = LGraphCanvas.node_colors.yellow.groupcolor
/* NOTE: this is not serialized and only there to make multiple
* note+ nodes in the same graph unique.
*/
uuid
/** Stores the dialog observer*/
resizeObserver
/** Live update the preview*/
live = true
/** DOM height by adding child size together*/
calculated_height = 0
/** ????*/
_raw_html
/** might not be needed anymore */
inner
/** the dialog DOM widget*/
dialog
/** widgets*/
/** used to store the raw value and display the parsed html at the same time*/
html_widget
/** hidden widgets for serialization*/
css_widget
edit_mode_widget
theme_widget
editorsContainer
/** ACE editors instances*/
html_editor
css_editor
constructor() { constructor() {
super() super()
this.uuid = shared.makeUUID() this.uuid = shared.makeUUID()
infoLogger('Constructing Note+ instance') infoLogger('Constructing Note+ instance')
shared.ensureMarkdownParser((_p) => {
this.updateHTML()
})
// - litegraph settings // - litegraph settings
this.collapsable = true this.collapsable = true
this.isVirtualNode = true this.isVirtualNode = true
@@ -159,35 +126,30 @@ class NotePlus extends LiteGraph.LGraphNode {
// - default values, serialization is done through widgets // - default values, serialization is done through widgets
this._raw_html = DEFAULT_MODE === 'html' ? DEFAULT_HTML : DEFAULT_MD this._raw_html = DEFAULT_MODE === 'html' ? DEFAULT_HTML : DEFAULT_MD
// - mardown converter
this.markdownConverter = new showdown.Converter({
tables: true,
strikethrough: true,
emoji: true,
ghCodeBlocks: true,
tasklists: true,
ghMentions: true,
smoothLivePreview: true,
simplifiedAutoLink: true,
parseImgDimensions: true,
openLinksInNewWindow: true,
})
// - state // - state
this.live = true this.live = true
this.calculated_height = 0 this.calculated_height = 0
// - add widgets // - add widgets
const inner = document.createElement('div') const cinner = document.createElement('div')
inner.style.margin = '0' this.inner = document.createElement('div')
inner.style.padding = '0'
inner.style.pointerEvents = 'none' cinner.append(this.inner)
this.html_widget = this.addDOMWidget('HTML', 'html', inner, { this.inner.classList.add('note-plus-preview')
cinner.style.margin = '0'
cinner.style.padding = '0'
this.html_widget = this.addDOMWidget('HTML', 'html', cinner, {
setValue: (val) => { setValue: (val) => {
this._raw_html = val this._raw_html = val
}, },
getValue: () => this._raw_html, getValue: () => this._raw_html,
getMinHeight: () => this.calculated_height, // (the edit button), getMinHeight: () => this.calculated_height, // (the edit button),
onDraw: () => {
// HACK: dirty hack for now until it's addressed upstream...
this.html_widget.element.style.pointerEvents = 'none'
// NOTE: not sure about this, it avoid the visual "bugs" but scrolling over the wrong area will affect zoom...
// this.html_widget.element.style.overflow = 'scroll'
},
hideOnZoom: false, hideOnZoom: false,
}) })
@@ -197,22 +159,48 @@ class NotePlus extends LiteGraph.LGraphNode {
} }
/** /**
* * @param {CanvasRenderingContext2D} ctx canvas context
* @param {CanvasRenderingContext2D} ctx * @param {any} _graphcanvas
* @param {LGraphCanvas} graphcanvas
* @returns
*/ */
onDrawForeground(ctx, _graphcanvas) { onDrawForeground(ctx, _graphcanvas) {
if (this.flags.collapsed) return if (this.flags.collapsed) return
this.drawEditIcon(ctx)
this.drawSideHandle(ctx)
// Define the size and position of the icon // DEBUG BACKGROUND
const iconSize = 14 // Size of the icon // ctx.fillStyle = 'rgba(0, 255, 0, 0.3)'
const iconMargin = 8 // Margin from the edges // const rect = this.rect
const x = this.size[0] - iconSize - iconMargin // ctx.fillRect(rect.x, rect.y, rect.width, rect.height)
const y = iconMargin * 1.5 }
drawSideHandle(ctx) {
const handleRect = this.sideHandleRect
const chamfer = 20
ctx.beginPath()
// top left
ctx.moveTo(handleRect.x, handleRect.y + chamfer)
// top right
ctx.lineTo(handleRect.x + handleRect.width, handleRect.y)
// bottom right
ctx.lineTo(
handleRect.x + handleRect.width,
handleRect.y + handleRect.height,
)
// bottom left
ctx.lineTo(handleRect.x, handleRect.y + handleRect.height - chamfer)
ctx.closePath()
ctx.fillStyle = 'rgba(255, 255, 255, 0.05)'
ctx.fill()
}
drawEditIcon(ctx) {
const rect = this.iconRect
// DEBUG ICON POSITION
// ctx.fillStyle = 'rgba(0, 255, 0, 0.3)'
// ctx.fillRect(rect.x, rect.y, rect.width, rect.height)
// Create a new Path2D object from SVG path data
const pencilPath = new Path2D( const pencilPath = new Path2D(
'M21.28 6.4l-9.54 9.54c-.95.95-3.77 1.39-4.4.76-.63-.63-.2-3.45.75-4.4l9.55-9.55a2.58 2.58 0 1 1 3.64 3.65z', 'M21.28 6.4l-9.54 9.54c-.95.95-3.77 1.39-4.4.76-.63-.63-.2-3.45.75-4.4l9.55-9.55a2.58 2.58 0 1 1 3.64 3.65z',
) )
@@ -220,41 +208,73 @@ class NotePlus extends LiteGraph.LGraphNode {
'M11 4H6a4 4 0 0 0-4 4v10a4 4 0 0 0 4 4h11c2.21 0 3-1.8 3-4v-5', 'M11 4H6a4 4 0 0 0-4 4v10a4 4 0 0 0 4 4h11c2.21 0 3-1.8 3-4v-5',
) )
// Draw the paths
ctx.save() ctx.save()
ctx.translate(x, y) // Position the icon on the canvas ctx.translate(rect.x, rect.y)
ctx.scale(iconSize / 32, iconSize / 32) // Scale the icon to the desired size ctx.scale(rect.width / 32, rect.height / 32)
ctx.strokeStyle = 'rgba(255,255,255,0.3)' ctx.strokeStyle = 'rgba(255,255,255,0.4)'
ctx.lineCap = 'round' ctx.lineCap = 'round'
ctx.lineJoin = 'round' ctx.lineJoin = 'round'
ctx.lineWidth = 2.4 ctx.lineWidth = 2.4
ctx.stroke(pencilPath) ctx.stroke(pencilPath)
ctx.stroke(folderPath) ctx.stroke(folderPath)
ctx.restore() ctx.restore()
} }
onMouseDown(_e, localPos, _graphcanvas) { /**
// Check if the click is within the pencil icon bounds * @param {number} x
const iconSize = 14 * @param {number} y
const iconMargin = 8 * @param {{x:number,y:number,width:number,height:number}} rect
const iconX = this.size[0] - iconSize - iconMargin * @returns {}
const iconY = iconMargin * 1.5 */
inRect(x, y, rect) {
if ( rect = rect || this.iconRect
localPos[0] > iconX && return (
localPos[0] < iconX + iconSize && x >= rect.x &&
localPos[1] > iconY && x <= rect.x + rect.width &&
localPos[1] < iconY + iconSize y >= rect.y &&
) { y <= rect.y + rect.height
// Pencil icon was clicked, open the editor )
this.openEditorDialog() }
return true // Return true to indicate the event was handled get rect() {
return {
x: 0,
y: 0,
width: this.size[0],
height: this.size[1],
} }
}
get sideHandleRect() {
const w = this.size[0]
const h = this.size[1]
return false // Return false to let the event propagate const bw = 32
const bho = 64
return {
x: w - bw,
y: bho,
width: bw,
height: h - bho * 1.5,
}
}
get iconRect() {
const iconSize = 32
const iconMargin = 16
return {
x: this.size[0] - iconSize - iconMargin,
y: iconMargin * 1.5,
width: iconSize,
height: iconSize,
}
}
onMouseDown(_e, localPos, _graphcanvas) {
if (this.inRect(localPos[0], localPos[1])) {
this.openEditorDialog()
return true
}
return false
} }
/* Hidden widgets to store note+ settings in the workflow (stripped in API)*/
setupSerializationWidgets() { setupSerializationWidgets() {
infoLogger('Setup Serializing widgets') infoLogger('Setup Serializing widgets')
@@ -283,15 +303,36 @@ class NotePlus extends LiteGraph.LGraphNode {
shared.hideWidgetForGood(this, this.css_widget) shared.hideWidgetForGood(this, this.css_widget)
shared.hideWidgetForGood(this, this.theme_widget) shared.hideWidgetForGood(this, this.theme_widget)
} }
setupDialog() { setupDialog() {
infoLogger('Setup dialog') infoLogger('Setup dialog')
// this.addWidget('button', 'Edit', 'Edit', this.openEditorDialog.bind(this))
this.dialog = new app.ui.dialog.constructor() this.dialog = new app.ui.dialog.constructor()
this.dialog.element.classList.add('comfy-settings') this.dialog.element.classList.add('comfy-settings')
Object.assign(this.dialog.element.style, {
position: 'absolute',
boxShadow: 'none',
})
const subcontainer = this.dialog.textElement.parentElement
if (subcontainer) {
Object.assign(subcontainer.style, {
width: '100%',
})
}
const closeButton = this.dialog.element.querySelector('button') const closeButton = this.dialog.element.querySelector('button')
closeButton.textContent = 'CANCEL' closeButton.textContent = 'CANCEL'
closeButton.id = 'cancel-editor-dialog'
closeButton.title =
"Cancel the changes since last opened (doesn't support live mode)"
closeButton.disabled = this.live
closeButton.style.background = this.live
? 'repeating-linear-gradient(45deg,#606dbc,#606dbc 10px,#465298 10px,#465298 20px)'
: ''
const saveButton = document.createElement('button') const saveButton = document.createElement('button')
saveButton.textContent = 'SAVE' saveButton.textContent = 'SAVE'
saveButton.onclick = () => { saveButton.onclick = () => {
@@ -313,32 +354,54 @@ class NotePlus extends LiteGraph.LGraphNode {
closeEditorDialog(accept) { closeEditorDialog(accept) {
infoLogger('Closing editor dialog', accept) infoLogger('Closing editor dialog', accept)
if (accept) { if (accept && !this.live) {
this.updateHTML(this.html_editor.getValue()) this.updateHTML(this.html_editor.getValue())
this.updateCSS(this.css_editor.getValue()) this.updateCSS(this.css_editor.getValue())
} }
if (this.resizeObserver) {
this.resizeObserver.disconnect()
this.resizeObserver = null
}
this.teardownEditors() this.teardownEditors()
this.dialog.close() this.dialog.close()
} }
/**
* @param {HTMLElement} elem
*/
hookResize(elem) {
if (!this.resizeObserver) {
const observer = () => {
this.html_editor.resize()
this.css_editor.resize()
Object.assign(this.editorsContainer.style, {
minHeight: `${(this.dialog.element.clientHeight / 100) * 50}px`, //'200px',
})
}
this.resizeObserver = new ResizeObserver(observer).observe(elem)
}
}
openEditorDialog() { openEditorDialog() {
infoLogger(`Current edit mode ${this.edit_mode_widget.value}`) infoLogger(`Current edit mode ${this.edit_mode_widget.value}`)
this.hookResize(this.dialog.element)
const container = document.createElement('div') const container = document.createElement('div')
Object.assign(container.style, { Object.assign(container.style, {
display: 'flex', display: 'flex',
gap: '10px', gap: '10px',
flexDirection: 'column', flexDirection: 'column',
}) })
const editorsContainer = document.createElement('div') this.editorsContainer = document.createElement('div')
Object.assign(editorsContainer.style, {
Object.assign(this.editorsContainer.style, {
display: 'flex', display: 'flex',
gap: '10px', gap: '10px',
flexDirection: 'row', flexDirection: 'row',
minHeight: this.dialog.element.offsetHeight, //'200px',
width: '100%',
}) })
container.append(editorsContainer) container.append(this.editorsContainer)
this.dialog.show('') this.dialog.show('')
this.dialog.textElement.append(container) this.dialog.textElement.append(container)
@@ -346,30 +409,39 @@ class NotePlus extends LiteGraph.LGraphNode {
const aceHTML = document.createElement('div') const aceHTML = document.createElement('div')
aceHTML.id = 'noteplus-html-editor' aceHTML.id = 'noteplus-html-editor'
Object.assign(aceHTML.style, { Object.assign(aceHTML.style, {
width: '300px', width: '100%',
height: '300px', height: '100%',
// backgroundColor: 'rgb(30,30,30)',
// color: 'whitesmoke', minWidth: '300px',
minHeight: 'inherit',
}) })
editorsContainer.append(aceHTML) this.editorsContainer.append(aceHTML)
const aceCSS = document.createElement('div') const aceCSS = document.createElement('div')
aceCSS.id = 'noteplus-css-editor' aceCSS.id = 'noteplus-css-editor'
Object.assign(aceCSS.style, { Object.assign(aceCSS.style, {
width: '300px', width: '100%',
height: '300px', height: '100%',
// backgroundColor: 'rgb(30,30,30)', minHeight: 'inherit',
// color: 'whitesmoke',
}) })
editorsContainer.append(aceCSS) this.editorsContainer.append(aceCSS)
const live_edit = document.createElement('input') const live_edit = document.createElement('input')
live_edit.type = 'checkbox' live_edit.type = 'checkbox'
live_edit.checked = this.live live_edit.checked = this.live
live_edit.onchange = () => { live_edit.onchange = () => {
this.live = live_edit.checked this.live = live_edit.checked
const cancel_button = this.dialog.element.querySelector(
'#cancel-editor-dialog',
)
if (cancel_button) {
cancel_button.disabled = this.live
cancel_button.style.background = this.live
? 'repeating-linear-gradient(45deg,#606dbc,#606dbc 10px,#465298 10px,#465298 20px)'
: ''
}
} }
//- "Dynamic" elements //- "Dynamic" elements
@@ -388,15 +460,14 @@ class NotePlus extends LiteGraph.LGraphNode {
const md = this.html_editor.getValue() const md = this.html_editor.getValue()
this.edit_mode_widget.value = 'html' this.edit_mode_widget.value = 'html'
select_mode.value = 'html' select_mode.value = 'html'
const html = this.markdownConverter.makeHtml(md) MTB.mdParser.parse(md).then((content) => {
this.html_widget.value = html this.html_widget.value = content
this.html_editor.setValue(html) this.html_editor.setValue(content)
this.html_editor.session.setMode('ace/mode/html') this.html_editor.session.setMode('ace/mode/html')
this.updateHTML(this.html_widget.value) this.updateHTML(this.html_widget.value)
convert_to_html.remove()
convert_to_html.remove() })
} }
firstButton.before(convert_to_html) firstButton.before(convert_to_html)
} }
} else { } else {
@@ -406,6 +477,19 @@ class NotePlus extends LiteGraph.LGraphNode {
} }
} }
select_mode.value = this.edit_mode_widget.value select_mode.value = this.edit_mode_widget.value
// the header for dragging the dialog
const header = document.createElement('div')
header.style.padding = '8px'
header.style.cursor = 'move'
header.style.backgroundColor = 'rgba(0,0,0,0.5)'
header.style.userSelect = 'none'
header.style.borderBottom = '1px solid #ddd'
header.textContent = 'MTB Note+ Editor'
container.prepend(header)
makeDraggable(this.dialog.element, header)
makeResizable(this.dialog.element)
} }
//- combobox //- combobox
let theme_select = this.dialog.element.querySelector('#theme_select') let theme_select = this.dialog.element.querySelector('#theme_select')
@@ -421,7 +505,7 @@ class NotePlus extends LiteGraph.LGraphNode {
option.textContent = label option.textContent = label
theme_select.append(option) theme_select.append(option)
} }
for (const t of themes) { for (const t of THEMES) {
addOption(t) addOption(t)
} }
@@ -491,54 +575,59 @@ class NotePlus extends LiteGraph.LGraphNode {
onCreate() { onCreate() {
errorLogger('NotePlus onCreate') errorLogger('NotePlus onCreate')
} }
configure(info) { restoreNodeState(info) {
super.configure(info)
infoLogger('Restoring serialized values', info)
// - update view from serialzed data
this.html_widget.element.id = `note-plus-${this.uuid}` this.html_widget.element.id = `note-plus-${this.uuid}`
this.setMode(this.edit_mode_widget.value) this.setMode(this.edit_mode_widget.value)
this.setTheme(this.theme_widget.value) this.setTheme(this.theme_widget.value)
this.updateHTML(this.html_widget.value) this.updateHTML(this.html_widget.value)
this.updateCSS(this.css_widget.value) this.updateCSS(this.css_widget.value)
this.setSize(info.size) if (info?.size) {
this.setSize(info.size)
}
}
configure(info) {
super.configure(info)
infoLogger('Restoring serialized values', info)
this.restoreNodeState(info)
// - update view from serialzed data
} }
onNodeCreated() { onNodeCreated() {
infoLogger('Node created', this.uuid) infoLogger('Node created', this.uuid)
this.html_widget.element.id = `note-plus-${this.uuid}` this.restoreNodeState({})
this.setMode(this.edit_mode_widget.value) // this.html_widget.element.id = `note-plus-${this.uuid}`
this.setTheme(this.theme_widget.value) // this.setMode(this.edit_mode_widget.value)
this.updateHTML(this.html_widget.value) // widget is populated here since we called super // this.setTheme(this.theme_widget.value)
this.updateCSS(this.css_widget.value) // this.updateHTML(this.html_widget.value) // widget is populated here since we called super
// this.updateCSS(this.css_widget.value)
} }
onRemoved() { onRemoved() {
infoLogger('Node removed', this.uuid) infoLogger('Node removed', this.uuid)
} }
getExtraMenuOptions() { getExtraMenuOptions() {
const options = [] const currentMode = this.edit_mode_widget.value
// { const newMode = currentMode === 'html' ? 'markdown' : 'html'
// content: string;
// callback?: ContextMenuEventListener;
// /** Used as innerHTML for extra child element */
// title?: string;
// disabled?: boolean;
// has_submenu?: boolean;
// submenu?: {
// options: ContextMenuItem[];
// } & IContextMenuOptions;
// className?: string;
// }
options.push({
content: `Set to ${
this.edit_mode_widget.value === 'html' ? 'markdown' : 'html'
}`,
callback: () => {
this.edit_mode_widget.value =
this.edit_mode_widget.value === 'html' ? 'markdown' : 'html'
this.updateHTML(this.html_widget.value)
},
})
return options const debugItems = window.MTB?.DEBUG
? [
{
content: 'Replace with demo content (debug)',
callback: () => {
this.html_widget.value = DEMO_CONTENT
},
},
]
: []
return [
...debugItems,
{
content: `Set to ${newMode}`,
callback: () => {
this.edit_mode_widget.value = newMode
this.updateHTML(this.html_widget.value)
},
},
]
} }
_setupEditor(editor) { _setupEditor(editor) {
@@ -663,17 +752,44 @@ class NotePlus extends LiteGraph.LGraphNode {
// this.setSize(this.computeSize()) // this.setSize(this.computeSize())
} }
updateHTML(val) { parserInitiated() {
const cleanHTML = DOMPurify.sanitize(val, { ADD_TAGS: ['iframe'] }) if (window.MTB?.mdParser) return true
this.html_widget.value = cleanHTML return false
}
// update our widget preview /** to easilty swap purification methods*/
if (this.edit_mode_widget.value === 'html') { purify(content) {
this.html_widget.element.innerHTML = cleanHTML return DOMPurify.sanitize(content, {
} else if (this.edit_mode_widget.value === 'markdown') { ADD_TAGS: ['iframe', 'detail', 'summary'],
this.html_widget.element.innerHTML = })
this.markdownConverter.makeHtml(cleanHTML) }
updateHTML(val) {
if (!this.parserInitiated()) {
return
} }
val = val || this.html_widget.value
const isHTML = this.edit_mode_widget.value === 'html'
const cleanHTML = this.purify(val)
const value = isHTML
? cleanHTML
: cleanHTML.replaceAll('&gt;', '>').replaceAll('&lt;', '<')
// .replaceAll('&amp;', '&')
// .replaceAll('&quot;', '"')
// .replaceAll('&#039;', "'")
this.html_widget.value = value
if (isHTML) {
this.inner.innerHTML = value
} else {
MTB.mdParser.parse(value).then((e) => {
this.inner.innerHTML = e
})
}
// this.html_widget.element.innerHTML = `<div id="note-plus-spacer"></div>${value}`
this.calculateHeight() this.calculateHeight()
// this.setSize(this.computeSize()) // this.setSize(this.computeSize())
} }
@@ -681,6 +797,27 @@ class NotePlus extends LiteGraph.LGraphNode {
app.registerExtension({ app.registerExtension({
name: 'mtb.noteplus', name: 'mtb.noteplus',
setup: () => {
app.ui.settings.addSetting({
id: 'mtb.noteplus.use-shiki',
category: ['mtb', 'Note+', 'use-shiki'],
name: 'Use shiki to highlight code',
tooltip:
'This will load a larger version of @mtb/markdown-parser that bundles shiki, it supports all shiki transformers (supported langs: html,css,python,markdown)',
type: 'boolean',
defaultValue: false,
attrs: {
style: {
// fontFamily: 'monospace',
},
},
async onChange(value) {
storage.set('np-use-shiki', value)
useShiki = value
},
})
},
registerCustomNodes() { registerCustomNodes() {
LiteGraph.registerNodeType('Note Plus (mtb)', NotePlus) LiteGraph.registerNodeType('Note Plus (mtb)', NotePlus)
File diff suppressed because one or more lines are too long
File diff suppressed because one or more lines are too long
+1 -1
Submodule wiki updated: 4db733ae92...fa7fec28a3