Compare commits

...
8 Commits
Author SHA1 Message Date
Mel Massadian d2da2949b4 feat: ✨ improve the I/O sidebar
- better options (sort, count)
- uses the new toast api instead of MTB.notify
2024-11-20 22:41:04 +01:00
Mel Massadian 9120d3ec42 chore: 🧹 add worktree to gitignores
for the experimental doc site at:
https://melmass.github.io/comfy_mtb/
2024-11-20 22:40:16 +01:00
Mel Massadian 2c8b7d790d feat: ✨ add UpscaleBBoxBy 2024-11-20 22:40:16 +01:00
Mel Massadian 540c8c9fa9 chore 🧹: add deprecations and experimental 2024-11-20 22:40:16 +01:00
Mel Massadian 88d51e5774 chore: 🧹 remove dupe code 2024-11-20 22:40:16 +01:00
Mel Massadian 8f47810b79 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:40:16 +01:00
Mel Massadian ae1ef0f914 feat: ✨ add Interpolate Condition 2024-11-20 22:40:06 +01:00
Mel Massadian b0e234b7ee feat: ✨ dump of wip things... 2024-11-20 22:39:23 +01:00
17 changed files with 1628 additions and 474 deletions
+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
+170 -36
View File
@@ -11,6 +11,8 @@ __version__ = "0.1.6"
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,20 @@ 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 .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 +55,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 +71,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 +123,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 +141,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 +161,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 +177,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 +190,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,19 +216,27 @@ 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"):
img_cache = None
with contextlib.suppress(ImportError):
from cachetools import TTLCache
img_cache = TTLCache(maxsize=100, ttl=5) # 1 min TTL
restore_deps = ["basicsr"] restore_deps = ["basicsr"]
onnx_deps = ["onnxruntime"] onnx_deps = ["onnxruntime"]
swap_deps = ["insightface"] + onnx_deps swap_deps = ["insightface"] + onnx_deps
@@ -307,8 +316,8 @@ if hasattr(PromptServer, "instance"):
) )
@PromptServer.instance.routes.post("/mtb/debug") @PromptServer.instance.routes.post("/mtb/debug")
async def set_debug(request): async def set_debug(request: Request):
json_data = await request.json() json_data: dict[str, bool] = await request.json()
enabled = json_data.get("enabled") enabled = json_data.get("enabled")
if enabled: if enabled:
os.environ["MTB_DEBUG"] = "true" os.environ["MTB_DEBUG"] = "true"
@@ -317,7 +326,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,10 +334,10 @@ 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
@@ -347,11 +356,136 @@ 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!"})
import asyncio
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:
if preview_params:
img = process_preview(img, preview_params)
if channel:
img = process_channel(img, channel)
if img_cache:
img_cache[cache_key] = img.getvalue()
return img_cache[cache_key]
return img.getvalue()
def process_preview(img: 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)
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}"'},
)
@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/debug") @PromptServer.instance.routes.get("/mtb/debug")
async def get_debug(request): async def get_debug(request: Request):
from . import endpoint from . import endpoint
reload(endpoint) _ = reload(endpoint)
enabled = "MTB_DEBUG" in os.environ enabled = "MTB_DEBUG" in os.environ
# 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", ""):
@@ -368,7 +502,7 @@ if hasattr(PromptServer, "instance"):
return web.json_response({"enabled": enabled}) return web.json_response({"enabled": enabled})
@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 +516,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)
+97 -19
View File
@@ -1,4 +1,9 @@
import csv import csv
import os
import secrets
import sys
from pathlib import Path
from typing import Any
from aiohttp import web from aiohttp import web
@@ -6,6 +11,8 @@ from .log import mklog
from .utils import ( from .utils import (
backup_file, backup_file,
import_install, import_install,
input_dir,
output_dir,
reqs_map, reqs_map,
run_command, run_command,
styles_dir, styles_dir,
@@ -14,15 +21,14 @@ 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=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]
@@ -50,6 +56,63 @@ def ACTIONS_installDependency(dependency_names=None):
# break # break
def ACTIONS_getUserImages(
mode: str,
count=200,
offset=0,
sort: str | None = None,
include_subfolders: bool = False,
):
# TODO: find a better name :s
enabled = "MTB_EXPOSE" in os.environ
if not enabled:
return {"error": "Session not authorized to getInputs"}
imgs = {}
entry_dir = input_dir if mode == "input" else output_dir
pattern = "**/*.png" if include_subfolders else "*.png"
entry_gen = entry_dir.glob(pattern)
entries = {}
if sort:
sort = sort.lower()
if sort == "none":
entries = entry_gen
elif sort == "modified":
entries = sorted(
entry_gen, key=lambda x: x.stat().st_mtime, reverse=True
)
elif sort == "modified-reverse":
entries = sorted(entry_gen, key=lambda x: x.stat().st_mtime)
elif sort == "name":
entries = sorted(entry_gen, key=lambda x: x.name)
elif sort == "name-reverse":
entries = sorted(entry_gen, key=lambda x: x.name, reverse=True)
else:
endlog.warning(f"Sort mode {sort} not supported")
entries = entry_gen
else:
entries = entry_gen
for i, img in enumerate(entries):
if i < offset:
continue
subfolder = (
img.parent.relative_to(entry_dir) if include_subfolders else ""
)
imgs[img.stem] = (
f"/mtb/view?filename={img.name}&width=512&type={mode}&subfolder="
f"{subfolder}"
f"&preview=&rand={secrets.randbelow(424242)}"
)
if i >= count + offset - 1:
break
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 +160,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 +172,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 +198,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 +224,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 +286,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 +305,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 +326,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 +338,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 +375,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 +416,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 +433,6 @@ def render_base_template(title, content):
<!-- Shared footer content here --> <!-- Shared footer content here -->
</footer> </footer>
</body> </body>
</html> </html>
""" """
+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,
+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,
+1
View File
@@ -8,3 +8,4 @@ rich
rich_argparse rich_argparse
matplotlib matplotlib
pillow pillow
cachetools
+6 -6
View File
@@ -163,9 +163,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 +175,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 +185,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 +198,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(
-33
View File
@@ -631,39 +631,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
+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)
} }
} }
}) })
}, },
}) })
+229
View File
@@ -0,0 +1,229 @@
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'
let offset = 0
let currentWidth = 200
let currentMode = 'input'
let currentSort = 'None'
const IMAGE_NODES = ['LoadImage']
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()
}
}
}
const getImgsFromUrls = (urls, target) => {
const imgs = []
if (urls === undefined) {
return imgs
}
for (const [key, url] of Object.entries(urls)) {
const a = makeElement('img')
a.src = url
a.width = currentWidth
if (currentMode === 'input') {
a.onclick = (_e) => {
const selected = app.canvas.selected_nodes
if (selected && Object.keys(selected).length === 0) {
app.extensionManager.toast.add({
severity: 'warn',
summary: 'No LoadImage 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}.png`)
}
}
} else {
a.onclick = (_e) =>
// window.MTB?.notify?.("Output import isn't supported yet...", 5000)
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,
})
}
imgs.push(a)
}
if (target !== undefined) {
target.append(...imgs)
}
return imgs
}
const getUrls = async () => {
const count = await api.getSetting('mtb.io-sidebar.count')
console.log('Sidebar count', count)
const inputs = await api.fetchApi('/mtb/actions', {
method: 'POST',
body: JSON.stringify({
name: 'getUserImages',
// mode, count, offset
args: [currentMode, count, offset, currentSort],
}),
})
const output = await inputs.json()
return output?.result || {}
}
if (window?.__COMFYUI_FRONTEND_VERSION__) {
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 urls = await getUrls(currentMode)
let imgs = {}
const cont = makeElement('div.mtb_sidebar')
const imgGrid = makeElement('div.mtb_img_grid')
const selector = makeSelect(['input', 'output'], currentMode)
selector.addEventListener('change', async (e) => {
const newMode = e.target.value
const changed = newMode !== currentMode
currentMode = newMode
if (changed) {
imgGrid.innerHTML = ''
const urls = await getUrls()
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()
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
}
},
})
}
+25
View File
@@ -0,0 +1,25 @@
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
}
+65 -60
View File
@@ -669,28 +669,29 @@ const mtb_widgets = {
} }
}, },
setup: () => { setup: () => {
app.ui.settings.addSetting({ app.ui.settings.addSetting({
id: 'mtb.Debug.enabled', id: 'mtb.Main.debug-enabled',
name: '[⚡mtb] Enable Debug (py and js)', category: ['mtb', 'Main', 'debug-enabled'],
type: 'boolean', name: 'Enable Debug (py and js)',
defaultValue: false, type: 'boolean',
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) {
if (!window.MTB) { if (!window.MTB) {
window.MTB = {} window.MTB = {}
} }
window.MTB.DEBUG = value window.MTB.DEBUG = value
if (value) { if (value) {
infoLogger('Enabled DEBUG mode') infoLogger('Enabled DEBUG mode')
} }
await api await api
.fetchApi('/mtb/debug', { .fetchApi('/mtb/debug', {
@@ -1110,45 +1111,49 @@ const mtb_widgets = {
break break
} }
//NOTE: dynamic nodes //NOTE: dynamic nodes
case 'Apply Text Template (mtb)': { case 'Apply Text Template (mtb)': {
shared.setupDynamicConnections(nodeType, 'var', '*') shared.setupDynamicConnections(nodeType, 'var', '*')
break break
} }
case 'Save Data Bundle (mtb)': { case 'Save Data Bundle (mtb)': {
shared.setupDynamicConnections(nodeType, 'data', '*') // [MASK,IMAGE] shared.setupDynamicConnections(nodeType, 'data', '*') // [MASK,IMAGE]
break break
} }
case 'Add To Playlist (mtb)': { case 'Add To Playlist (mtb)': {
shared.setupDynamicConnections(nodeType, 'video', 'VIDEO') shared.setupDynamicConnections(nodeType, 'video', 'VIDEO')
break break
} }
case 'Psd Save (mtb)': { case 'Interpolate Condition (mtb)': {
shared.setupDynamicConnections(nodeType, 'input_', 'PSDLAYER') shared.setupDynamicConnections(nodeType, 'condition', 'CONDITIONING')
break break
} }
// case 'Text Encode Frames (mtb)' : { case 'Psd Save (mtb)': {
// shared.setupDynamicConnections(nodeType, 'input_', 'IMAGE') shared.setupDynamicConnections(nodeType, 'input_', 'PSDLAYER')
// break break
// } }
case 'Stack Images (mtb)': // case 'Text Encode Frames (mtb)' : {
case 'Concat Images (mtb)': { // shared.setupDynamicConnections(nodeType, 'input_', 'IMAGE')
shared.setupDynamicConnections(nodeType, 'image', 'IMAGE') // break
break // }
} case 'Stack Images (mtb)':
case 'Audio Sequence (mtb)': case 'Concat Images (mtb)': {
case 'Audio Stack (mtb)': { shared.setupDynamicConnections(nodeType, 'image', 'IMAGE')
shared.setupDynamicConnections(nodeType, 'audio', 'AUDIO') break
break }
} case 'Audio Sequence (mtb)':
case 'Batch Float Assemble (mtb)': case 'Audio Stack (mtb)': {
case 'Batch Float Math (mtb)': shared.setupDynamicConnections(nodeType, 'audio', 'AUDIO')
case 'Plot Batch Float (mtb)': { break
shared.setupDynamicConnections(nodeType, 'floats', 'FLOATS') }
break case 'Batch Float Assemble (mtb)':
} case 'Batch Float Math (mtb)':
case 'Batch Merge (mtb)': { case 'Plot Batch Float (mtb)': {
shared.setupDynamicConnections(nodeType, 'batches', 'IMAGE') shared.setupDynamicConnections(nodeType, 'floats', 'FLOATS')
break
}
case 'Batch Merge (mtb)': {
shared.setupDynamicConnections(nodeType, 'batches', 'IMAGE')
break break
} }