Compare commits

..
5 Commits
Author SHA1 Message Date
Mel Massadian d576190f07 fix: 🐛 category for settings 2024-11-20 21:35:07 +01:00
Mel Massadian 247ad12ec3 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-18 01:53:04 +01:00
Mel Massadian 881573e2f2 feat: ✨ use the new parser for documentations
- might also fix #210
2024-11-18 01:32:29 +01:00
Mel Massadian 6f3a5b5c71 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-18 00:48:04 +01:00
Mel Massadian fa017fbb81 chore: 🧹 update externs
- remove showdown
- update dompurify
2024-11-18 00:45:53 +01:00
29 changed files with 1808 additions and 4294 deletions
-2
View File
@@ -12,8 +12,6 @@ jobs:
steps:
- name: ♻️ Check out code
uses: actions/checkout@v4
with:
submodules: true
- name: 📦 Publish Custom Node
uses: Comfy-Org/publish-node-action@main
with:
-3
View File
@@ -6,6 +6,3 @@ node_modules/
compose.yaml
comfy_mtb.wsb
Dockerfile
# I store the gh-pages worktrees (src & build) there
.worktrees
+57 -209
View File
@@ -7,12 +7,10 @@
#
###
__version__ = "0.2.1"
__version__ = "0.1.6"
import os
from aiohttp.web_request import Request
# TODO: don't override this if the user has that setup already
if not os.environ.get("TF_FORCE_GPU_ALLOW_GROWTH"):
os.environ["TF_FORCE_GPU_ALLOW_GROWTH"] = "true"
@@ -33,21 +31,24 @@ from pathlib import Path
from aiohttp import web
from server import PromptServer
import nodes
from .endpoint import endlog
from .install import get_node_dependencies
from .log import blue_text, cyan_text, get_label, get_summary, log
from .utils import comfy_dir, here
NODE_CLASS_MAPPINGS: dict[str, type] = {}
NODE_DISPLAY_NAME_MAPPINGS: dict[str, str] = {}
NODE_CLASS_MAPPINGS_DEBUG: dict[str, str | None] = {}
NODE_CLASS_MAPPINGS = {}
NODE_DISPLAY_NAME_MAPPINGS = {}
NODE_CLASS_MAPPINGS_DEBUG = {}
WEB_DIRECTORY = "./web"
def extract_nodes_from_source(filename: Path):
source_code = ""
source_code = filename.read_text(encoding="utf-8")
nodes: list[str] = []
nodes = []
try:
parsed = ast.parse(source_code)
@@ -56,15 +57,14 @@ def extract_nodes_from_source(filename: Path):
target = node.targets[0]
if isinstance(target, ast.Name) and target.id == "__nodes__":
value = ast.get_source_segment(source_code, node.value)
if value:
node_value = ast.parse(value).body[0].value
if isinstance(node_value, ast.List | ast.Tuple):
nodes.extend(
str(element.id)
for element in node_value.elts
if isinstance(element, ast.Name)
)
break
node_value = ast.parse(value).body[0].value
if isinstance(node_value, (ast.List, ast.Tuple)):
nodes.extend(
element.id
for element in node_value.elts
if isinstance(element, ast.Name)
)
break
except SyntaxError:
log.error("Failed to parse")
return nodes
@@ -72,8 +72,8 @@ def extract_nodes_from_source(filename: Path):
def load_nodes():
errors: list[str] = []
nodes: list[type] = []
nodes_failed: list[str] = []
nodes = []
nodes_failed = []
for filename in (here / "nodes").iterdir():
if filename.suffix == ".py":
@@ -124,8 +124,7 @@ def uninstall_old_web_extensions():
shutil.rmtree(web_mtb)
except Exception as e:
log.warning(
f"""Failed to remove web mtb directory: {e}
Please manually remove it from disk ({web_mtb}) and restart the server."""
f"Failed to remove web mtb directory: {e}\nPlease manually remove it from disk ({web_mtb}) and restart the server."
)
@@ -142,7 +141,7 @@ def wiki_to_classname(s: str):
def classname_to_wiki(s: str):
classname = s.replace("MTB_", "")
parts: list[str] = []
parts = []
start = 0
for i in range(1, len(classname)):
if classname[i].isupper():
@@ -162,6 +161,8 @@ if wiki.exists() and wiki.is_dir():
# - REGISTER NODES
MTB_EXPORT = os.environ.get("MTB_EXPORT")
nodes, failed = load_nodes()
@@ -178,7 +179,7 @@ for node_class in nodes:
node_class.DESCRIPTION = node_class.__doc__
if MTB_EXPORT:
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"
)
@@ -191,15 +192,12 @@ for node_class in nodes:
NODE_CLASS_MAPPINGS[node_label] = node_class
NODE_DISPLAY_NAME_MAPPINGS[class_name] = node_label
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"):
with open(here / "node_list.json", "w") as f:
_ = f.write(
f.write(
json.dumps(
{
k: NODE_CLASS_MAPPINGS_DEBUG[k]
@@ -217,31 +215,29 @@ log.debug(
)
)
log.info(f"loaded {cyan_text(str(len(nodes)))} nodes successfuly")
log.info(f"loaded {cyan_text(len(nodes))} nodes successfuly")
if failed:
with contextlib.suppress(Exception):
base_url, port = utils.get_server_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."
)
log.debug(failed)
# - ENDPOINT
if hasattr(PromptServer, "instance"):
img_cache = None
prompt_cache = None
with contextlib.suppress(ImportError):
from cachetools import TTLCache
img_cache = TTLCache(maxsize=100, ttl=5) # 1 min TTL
prompt_cache = TTLCache(maxsize=100, ttl=5) # 1 min TTL
node_dependency_mapping = get_node_dependencies()
restore_deps = ["basicsr"]
onnx_deps = ["onnxruntime"]
swap_deps = ["insightface"] + onnx_deps
node_dependency_mapping = {
"QrCode": ["qrcode"],
"DeepBump": onnx_deps,
"FaceSwap": swap_deps,
"LoadFaceSwapModel": swap_deps,
"LoadFaceAnalysisModel": restore_deps,
}
PromptServer.instance.app.router.add_static(
"/mtb-assets/", path=(here / "html").as_posix()
@@ -310,10 +306,10 @@ if hasattr(PromptServer, "instance"):
}
)
@PromptServer.instance.routes.post("/mtb/server-info")
async def set_server_info(request: Request):
json_data: dict[str, bool] = await request.json()
enabled = json_data.get("debug")
@PromptServer.instance.routes.post("/mtb/debug")
async def set_debug(request):
json_data = await request.json()
enabled = json_data.get("enabled")
if enabled:
os.environ["MTB_DEBUG"] = "true"
log.setLevel(logging.DEBUG)
@@ -321,7 +317,7 @@ if hasattr(PromptServer, "instance"):
elif "MTB_DEBUG" in os.environ:
# del os.environ["MTB_DEBUG"]
_ = os.environ.pop("MTB_DEBUG")
os.environ.pop("MTB_DEBUG")
log.setLevel(logging.INFO)
return web.json_response(
@@ -329,17 +325,17 @@ if hasattr(PromptServer, "instance"):
)
@PromptServer.instance.routes.get("/mtb")
async def get_home(request: Request):
async def get_home(request):
from . import endpoint
_ = reload(endpoint)
reload(endpoint)
# Check if the request prefers HTML content
if "text/html" in request.headers.get("Accept", ""):
# # Return an HTML page
html_response = """
<div class="flex-container menu">
<a href="/mtb/manage">manage</a>
<a href="/mtb/server-info">Server Info</a>
<a href="/mtb/debug">debug</a>
<a href="/mtb/status">status</a>
</div>
"""
@@ -351,176 +347,28 @@ if hasattr(PromptServer, "instance"):
# Return JSON for other requests
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:
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):
@PromptServer.instance.routes.get("/mtb/debug")
async def get_debug(request):
from . import endpoint
_ = reload(endpoint)
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>"""
reload(endpoint)
enabled = "MTB_DEBUG" in os.environ
# Check if the request prefers HTML content
if "text/html" in request.headers.get("Accept", ""):
# # Return an HTML page
html_response = ""
html_response += render_property(
"Debug", "Enabled" if isdebug else "Disabled"
)
html_response += render_property("Exposed", str(exposed))
html_response = f"""
<h1>MTB Debug Status: {'Enabled' if enabled else 'Disabled'}</h1>
"""
return web.Response(
text=endpoint.render_base_template(
"Server Info", html_response
),
text=endpoint.render_base_template("Debug", html_response),
content_type="text/html",
)
# Return JSON for other requests
return web.json_response({"exposed": exposed, "debug": isdebug})
return web.json_response({"enabled": enabled})
@PromptServer.instance.routes.get("/mtb/actions")
async def no_route(request: Request):
async def no_route(request):
from . import endpoint
if "text/html" in request.headers.get("Accept", ""):
@@ -534,7 +382,7 @@ if hasattr(PromptServer, "instance"):
return web.json_response({"message": "actions has no get for now"})
@PromptServer.instance.routes.post("/mtb/actions")
async def do_action(request: Request):
async def do_action(request):
from . import endpoint
reload(endpoint)
+20 -152
View File
@@ -1,20 +1,10 @@
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 .install import get_node_dependencies
from .log import mklog
from .utils import (
SortMode,
backup_file,
build_glob_patterns,
glob_multiple,
import_install,
reqs_map,
run_command,
@@ -24,25 +14,18 @@ from .utils import (
endlog = mklog("mtb endpoint")
# - ACTIONS
import sys
from pathlib import Path
import_install("requirements")
def ACTIONS_installDependency(dependency_names: list[str] | None = None):
def ACTIONS_installDependency(dependency_names=None):
if dependency_names is None:
# return web.Response(text="No dependency name provided", status=400)
return {"error": "No dependency name provided"}
endlog.debug(f"Received Install Dependency request for {dependency_names}")
# reqs = []
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:
run_command(
[Path(sys.executable), "-m", "pip", "install"] + resolved_names
@@ -67,106 +50,6 @@ def ACTIONS_installDependency(dependency_names: list[str] | None = None):
# 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):
from .nodes.conditions import MTB_StylesLoader
@@ -214,7 +97,7 @@ def ACTIONS_saveStyle(data):
csv_writer.writerow(row)
async def do_action(request: web.Request) -> web.Response:
async def do_action(request) -> web.Response:
endlog.debug("Init action request")
request_data = await request.json()
name = request_data.get("name")
@@ -226,12 +109,7 @@ async def do_action(request: web.Request) -> web.Response:
method = globals().get(method_name)
if callable(method):
result = None
if args:
result = method(*args) if isinstance(args, list) else method(args)
else:
result = method()
result = method(args) if args else method()
endlog.debug(f"Action result: {result}")
return web.json_response({"result": result})
@@ -252,13 +130,10 @@ async def do_action(request: web.Request) -> web.Response:
# - HTML UTILS
def dependencies_button(name: str, dependencies: list[str]) -> str:
def dependencies_button(name, dependencies):
deps = ",".join([f"'{x}'" for x in dependencies])
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>
"""
@@ -278,7 +153,7 @@ def csv_editor():
html_out = """
<div id="style-editor">
<h1>Style Editor</h1>
"""
for current, styles in style_files.items():
current_out = f"<h3>{current}</h3>"
@@ -340,14 +215,11 @@ def render_tab_view(**kwargs):
"""
def add_foldable_region(title: str, content: str):
def add_foldable_region(title, content):
symbol_id = f"{title}-symbol"
return f"""
<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>
{title}
</div>
@@ -359,9 +231,7 @@ def add_foldable_region(title: str, content: str):
"""
def add_split_pane(
left_content: str, right_content: str, *, vertical: bool = True
):
def add_split_pane(left_content, right_content, vertical=True):
orientation = "vertical" if vertical else "horizontal"
return f"""
<div class="split-pane {orientation}">
@@ -380,7 +250,7 @@ def add_split_pane(
"""
def add_dropdown(title: str, options: list[str]):
def add_dropdown(title, options):
option_str = "\n".join(
[f"<option value='{opt}'>{opt}</option>" for opt in options]
)
@@ -392,13 +262,13 @@ def add_dropdown(title: str, options: list[str]):
"""
def render_table(table_dict: dict[str, Any], sort=True, title=None):
table_list = sorted(
def render_table(table_dict, sort=True, title=None):
table_dict = sorted(
table_dict.items(), key=lambda item: item[0]
) # Sort the dictionary by keys
table_rows = ""
for name, item in table_list:
for name, item in table_dict:
if isinstance(item, dict):
if "dependencies" in item:
table_rows += f"<tr><td>{name}</td><td>"
@@ -429,12 +299,12 @@ def render_table(table_dict: dict[str, Any], sort=True, title=None):
<tbody>
{table_rows}
</tbody>
</table>
</table>
</div>
"""
def render_base_template(title: str, content: str):
def render_base_template(title, content):
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"""
<!DOCTYPE html>
@@ -470,9 +340,7 @@ def render_base_template(title: str, content: str):
<header>
<a href="/">Back to Comfy</a>
<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>
<a style="width:128px;text-align:center" href="https://www.github.com/melmass/comfy_mtb">
{github_icon_svg}
@@ -487,6 +355,6 @@ def render_base_template(title: str, content: str):
<!-- Shared footer content here -->
</footer>
</body>
</html>
"""
+5 -17
View File
@@ -27,7 +27,7 @@ export def "comfy start" [--clean,--old-ui, --listen] {
let root = get_root --clean=($clean)
cd $root
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 {[]})
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 {[]})
}
# update comfy itself and merge master in current branch
@@ -67,14 +67,8 @@ export def "comfy update" [
git checkout master
print $"(ansi yellow_italic)Fetching and pulling remote updates(ansi reset)"
if ($clean) {
git fetch local master
git pull local master
} else {
git fetch
git pull
}
git fetch
git pull
print $"(ansi yellow_italic)Back to our branch \(($branch_name)\)(ansi reset)"
git checkout -
@@ -141,7 +135,7 @@ export def "comfy update_extensions" [--clean] {
let root = get_root --clean=($clean)
cd $root
cd custom_nodes
git multipull . -s -q
git multipull .
}
def --env path-add [pth] {
@@ -152,7 +146,7 @@ def --env path-add [pth] {
export-env {
$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
@@ -160,12 +154,6 @@ export-env {
$env.COMFY_CLEAN_ROOT = ($env.COMFY_ROOT | path dirname | path join ComfyClean)
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)
overlay use ../../.venv/Scripts/activate.nu
}
+28 -64
View File
@@ -43,28 +43,10 @@ pip_map = {
"tb-nightly": "tensorboard",
"protobuf": "google.protobuf",
"qrcode[pil]": "qrcode",
"requirements-parser": "requirements",
"requirements-parser": "requirements"
# 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
# region ansi
@@ -142,12 +124,12 @@ def print_formatted(text, *formats, color=None, background=None, **kwargs):
header = "[mtb install] "
# Handle console encoding for Unicode characters (utf-8)
encoded_header = header.encode(
sys.stdout.encoding, errors="replace"
).decode(sys.stdout.encoding)
encoded_text = formatted_text.encode(
sys.stdout.encoding, errors="replace"
).decode(sys.stdout.encoding)
encoded_header = header.encode(sys.stdout.encoding, errors="replace").decode(
sys.stdout.encoding
)
encoded_text = formatted_text.encode(sys.stdout.encoding, errors="replace").decode(
sys.stdout.encoding
)
print(
" " * len(encoded_header)
@@ -181,9 +163,7 @@ def run_command(cmd, ignored_lines_start=None):
try:
_run_command(shell_cmd, ignored_lines_start)
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)
except KeyboardInterrupt:
@@ -258,7 +238,7 @@ def suppress_std():
def get_local_version():
init_file = os.path.join(os.path.dirname(__file__), "__init__.py")
if os.path.isfile(init_file):
with open(init_file) as f:
with open(init_file, "r") as f:
tree = ast.parse(f.read())
for node in ast.walk(tree):
if isinstance(node, ast.Assign):
@@ -276,16 +256,13 @@ def download_file(url, file_name):
with requests.get(url, stream=True) as response:
response.raise_for_status()
total_size = int(response.headers.get("content-length", 0))
with (
open(file_name, "wb") as file,
tqdm(
desc=file_name.stem,
total=total_size,
unit="B",
unit_scale=True,
unit_divisor=1024,
) as progress_bar,
):
with open(file_name, "wb") as file, tqdm(
desc=file_name.stem,
total=total_size,
unit="B",
unit_scale=True,
unit_divisor=1024,
) as progress_bar:
for chunk in response.iter_content(chunk_size=8192):
file.write(chunk)
progress_bar.update(len(chunk))
@@ -325,9 +302,7 @@ def import_or_install(requirement, dry=False):
pip_install_name = pip_name + pip_spec
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:
print_formatted(
f"Dry-run: Package {pip_install_name} would be installed (import name: '{import_name}').",
@@ -335,9 +310,7 @@ def import_or_install(requirement, dry=False):
)
else:
try:
run_command(
[executable, "-m", "pip", "install", pip_install_name]
)
run_command([executable, "-m", "pip", "install", pip_install_name])
print_formatted(
f"Package {pip_install_name} installed successfully using pip package name (import name: '{import_name}')",
"bold",
@@ -353,9 +326,13 @@ def import_or_install(requirement, dry=False):
def get_github_assets(tag=None):
if tag:
tag_url = f"https://api.github.com/repos/{repo_owner}/{repo_name}/releases/tags/{tag}"
tag_url = (
f"https://api.github.com/repos/{repo_owner}/{repo_name}/releases/tags/{tag}"
)
else:
tag_url = f"https://api.github.com/repos/{repo_owner}/{repo_name}/releases/latest"
tag_url = (
f"https://api.github.com/repos/{repo_owner}/{repo_name}/releases/latest"
)
response = requests.get(tag_url)
if response.status_code == 404:
# print_formatted(
@@ -384,9 +361,7 @@ except ImportError:
def main():
if len(sys.argv) == 1:
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
if all(arg not in ("-p", "--path") for arg in sys.argv):
@@ -422,12 +397,8 @@ def main():
else:
repo_dir = clone_dir / repo_name
if not repo_dir.exists():
print_formatted(
f"Cloning to {repo_dir}...", "italic", color="yellow"
)
run_command(
["git", "clone", "--recursive", repo_url, repo_dir]
)
print_formatted(f"Cloning to {repo_dir}...", "italic", color="yellow")
run_command(["git", "clone", "--recursive", repo_url, repo_dir])
else:
print_formatted(
f"Directory {repo_dir} already exists, we will update it..."
@@ -438,14 +409,7 @@ def main():
print_formatted("Checking environment...", "italic", color="yellow")
missing_deps = []
install_cmd = [
executable,
"-m",
"pip",
"install",
"-r",
"requirements.txt",
]
install_cmd = [executable, "-m", "pip", "install", "-r", "requirements.txt"]
run_command(install_cmd)
print_formatted(
+10 -226
View File
@@ -1,5 +1,4 @@
from io import BytesIO
from typing import Literal
import cv2
import numpy as np
@@ -411,14 +410,7 @@ class MTB_BatchFloat:
RETURN_TYPES = ("FLOATS",)
CATEGORY = "mtb/batch"
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",
):
def set_floats(self, mode, count, min, max, easing):
if mode == "Steps" and count == 1:
raise ValueError(
"Steps mode requires at least a count of 2 values"
@@ -437,210 +429,6 @@ class MTB_BatchFloat:
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:
"""Merges multiple image batches with different frame counts"""
@@ -923,9 +711,7 @@ class MTB_PlotBatchFloat:
ax.set_xlim(1, max_length) # Set X-axis limits
np.random.seed(seed)
colors = np.random.rand(len(kwargs), 3) # Generate random RGB values
for color, (label, values) in zip(
colors, kwargs.items(), strict=False
):
for color, (label, values) in zip(colors, kwargs.items()):
ax.plot(x_values[: len(values)], values, label=label, color=color)
ax.legend(
title="Legend",
@@ -1240,19 +1026,17 @@ class MTB_BatchShake:
__nodes__ = [
MTB_Batch2dTransform,
MTB_BatchFloat,
MTB_Batch2dTransform,
MTB_BatchShape,
MTB_BatchMake,
MTB_BatchFloatAssemble,
MTB_BatchFloatFill,
MTB_BatchFloatNormalize,
MTB_BatchMerge,
MTB_BatchShake,
MTB_PlotBatchFloat,
MTB_BatchTimeWrap,
MTB_BatchFloatFit,
MTB_BatchFloatMath,
MTB_BatchFloatNormalize,
MTB_BatchMake,
MTB_BatchMerge,
MTB_BatchSequence,
MTB_BatchSequencePlus,
MTB_BatchShake,
MTB_BatchShape,
MTB_BatchTimeWrap,
MTB_PlotBatchFloat,
]
+1 -123
View File
@@ -3,127 +3,10 @@ import shutil
from pathlib import Path
import folder_paths
import torch
from ..log import log
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:
@classmethod
@@ -330,9 +213,4 @@ class MTB_StylesLoader:
return (self.options[style_name][0], self.options[style_name][1])
__nodes__ = [
MTB_SmartStep,
MTB_StylesLoader,
MTB_InterpolateClipSequential,
MTB_InterpolateCondition,
]
__nodes__ = [MTB_SmartStep, MTB_StylesLoader, MTB_InterpolateClipSequential]
+1 -38
View File
@@ -59,36 +59,6 @@ class MTB_SplitBbox:
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:
"""From a mask extract the bounding box"""
@@ -372,11 +342,4 @@ class MTB_Uncrop:
return (pil2tensor(out_images),)
__nodes__ = [
MTB_BboxFromMask,
MTB_Bbox,
MTB_Crop,
MTB_Uncrop,
MTB_SplitBbox,
MTB_UpscaleBboxBy,
]
__nodes__ = [MTB_BboxFromMask, MTB_Bbox, MTB_Crop, MTB_Uncrop, MTB_SplitBbox]
-2
View File
@@ -78,7 +78,6 @@ class MTB_LoadFaceEnhanceModel:
RETURN_NAMES = ("model",)
FUNCTION = "load_model"
CATEGORY = "mtb/facetools"
DEPRECATED = True
def load_model(self, model_name, upscale=2, bg_upsampler=None):
from gfpgan import GFPGANer
@@ -164,7 +163,6 @@ class MTB_RestoreFace:
RETURN_TYPES = ("IMAGE",)
FUNCTION = "restore"
CATEGORY = "mtb/facetools"
DEPRECATED = True
@classmethod
def INPUT_TYPES(cls):
-3
View File
@@ -40,7 +40,6 @@ class MTB_LoadFaceAnalysisModel:
RETURN_TYPES = ("FACE_ANALYSIS_MODEL",)
FUNCTION = "load_model"
CATEGORY = "mtb/facetools"
DEPRECATED = True
def load_model(self, faceswap_model: str):
if faceswap_model == "antelopev2":
@@ -78,7 +77,6 @@ class MTB_LoadFaceSwapModel:
RETURN_TYPES = ("FACESWAP_MODEL",)
FUNCTION = "load_model"
CATEGORY = "mtb/facetools"
DEPRECATED = True
def load_model(self, faceswap_model: str):
model_path = get_model_path("insightface", faceswap_model)
@@ -128,7 +126,6 @@ class MTB_FaceSwap:
RETURN_TYPES = ("IMAGE",)
FUNCTION = "swap"
CATEGORY = "mtb/facetools"
DEPRECATED = True
def swap(
self,
+4 -11
View File
@@ -1,4 +1,5 @@
from pathlib import Path
from typing import List
import comfy
import comfy.model_management as model_management
@@ -14,13 +15,10 @@ from ..utils import get_model_path
class MTB_LoadFilmModel:
"""Loads a FILM model
[DEPRECATED] Use ComfyUI-FrameInterpolation instead
"""
"""Loads a FILM model"""
@staticmethod
def get_models() -> list[Path]:
def get_models() -> List[Path]:
models_paths = get_model_path("FILM").iterdir()
return [x for x in models_paths if x.suffix in [".onnx", ".pth"]]
@@ -39,7 +37,6 @@ class MTB_LoadFilmModel:
RETURN_TYPES = ("FILM_MODEL",)
FUNCTION = "load_model"
CATEGORY = "mtb/frame iterpolation"
DEPRECATED = True
def load_model(self, film_model: str):
model_path = get_model_path("FILM", film_model)
@@ -59,10 +56,7 @@ class MTB_LoadFilmModel:
class MTB_FilmInterpolation:
"""Google Research FILM frame interpolation for large motion
[DEPRECATED] Use ComfyUI-FrameInterpolation instead
"""
"""Google Research FILM frame interpolation for large motion"""
@classmethod
def INPUT_TYPES(cls):
@@ -77,7 +71,6 @@ class MTB_FilmInterpolation:
RETURN_TYPES = ("IMAGE",)
FUNCTION = "do_interpolation"
CATEGORY = "mtb/frame iterpolation"
DEPRECATED = True
def do_interpolation(
self,
+3 -1
View File
@@ -702,6 +702,7 @@ class MTB_Blur:
)
blurred_images.append(blurred)
image_np = np.array(blurred_images)
else:
for i in range(image.size(0)):
blurred = gaussian(
@@ -709,7 +710,8 @@ class MTB_Blur:
)
blurred_images.append(blurred)
return (np2tensor(blurred_images),)
image_np = np.array(blurred_images)
return (np2tensor(image_np).squeeze(0),)
class MTB_Sharpen:
+25 -169
View File
@@ -1,11 +1,4 @@
import json
import os
import numpy as np
import torch
from comfy.cli_args import args
from PIL import Image
from PIL.PngImagePlugin import PngInfo
from ..log import log
@@ -15,21 +8,13 @@ class MTB_StackImages:
@classmethod
def INPUT_TYPES(cls):
return {
"required": {"vertical": ("BOOLEAN", {"default": False})},
"optional": {
"match_method": (
["error", "smallest", "largest"],
{"default": "error"},
)
},
}
return {"required": {"vertical": ("BOOLEAN", {"default": False})}}
RETURN_TYPES = ("IMAGE",)
FUNCTION = "stack"
CATEGORY = "mtb/image utils"
def stack(self, vertical, match_method="error", **kwargs):
def stack(self, vertical, **kwargs):
if not kwargs:
raise ValueError("At least one tensor must be provided.")
@@ -47,50 +32,23 @@ class MTB_StackImages:
self.duplicate_frames(tensor, max_batch_size)
for tensor in normalized_tensors
]
if match_method != "error":
if vertical:
# match widths
widths = [tensor.shape[2] for tensor in normalized_tensors]
target_width = (
min(widths) if match_method == "smallest" else max(widths)
)
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:
if vertical:
width = normalized_tensors[0].shape[2]
if any(
tensor.shape[2] != width for tensor in normalized_tensors
):
raise ValueError(
"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
if vertical:
width = normalized_tensors[0].shape[2]
if any(tensor.shape[2] != width for tensor in normalized_tensors):
raise ValueError(
"All tensors must have the same width "
"for vertical stacking."
)
dim = 1
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 = 2
stacked_tensor = torch.cat(normalized_tensors, dim=dim)
@@ -106,7 +64,7 @@ class MTB_StackImages:
elif channels == 3:
alpha_channel = torch.ones(
tensor.shape[:-1] + (1,), device=tensor.device
)
) # Add an alpha channel
return torch.cat((tensor, alpha_channel), dim=-1)
else:
raise ValueError(
@@ -129,30 +87,6 @@ class MTB_StackImages:
else:
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:
"""Pick a specific number of images from a batch.
@@ -179,6 +113,11 @@ class MTB_PickFromBatch:
# Limit count to the available number of images in the batch
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":
selected_tensors = image[-count:]
@@ -188,87 +127,4 @@ class MTB_PickFromBatch:
return (selected_tensors,)
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]
__nodes__ = [MTB_StackImages, MTB_PickFromBatch]
+19 -55
View File
@@ -2,9 +2,9 @@ import json
import subprocess
import uuid
from pathlib import Path
from typing import List, Optional
import comfy.model_management as model_management
import comfy.utils
import folder_paths
import numpy as np
import torch
@@ -41,7 +41,6 @@ class MTB_ReadPlaylist:
RETURN_TYPES = ("PLAYLIST",)
FUNCTION = "read_playlist"
CATEGORY = "mtb/IO"
EXPERIMENTAL = True
def read_playlist(
self,
@@ -84,7 +83,6 @@ class MTB_AddToPlaylist:
OUTPUT_NODE = True
FUNCTION = "add_to_playlist"
CATEGORY = "mtb/IO"
EXPERIMENTAL = True
def add_to_playlist(
self,
@@ -119,10 +117,7 @@ class MTB_AddToPlaylist:
class MTB_ExportWithFfmpeg:
"""Export with FFmpeg (Experimental).
[DEPRACATED] Use VHS nodes instead
"""
"""Export with FFmpeg (Experimental)"""
@classmethod
def INPUT_TYPES(cls):
@@ -148,7 +143,6 @@ class MTB_ExportWithFfmpeg:
RETURN_TYPES = ("VIDEO",)
OUTPUT_NODE = True
FUNCTION = "export_prores"
DEPRECATED = True
CATEGORY = "mtb/IO"
def export_prores(
@@ -157,9 +151,10 @@ class MTB_ExportWithFfmpeg:
prefix: str,
format: str,
codec: str,
images: torch.Tensor | None = None,
playlist: list[str] | None = None,
images: Optional[torch.Tensor] = None,
playlist: Optional[List[str]] = None,
):
pix_fmt = "rgb48le" if codec == "prores_ks" else "yuv420p"
file_ext = format
file_id = f"{prefix}_{uuid.uuid4()}.{file_ext}"
@@ -213,11 +208,9 @@ class MTB_ExportWithFfmpeg:
frames = tensor2np(images)
log.debug(f"Frames type {type(frames[0])}")
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":
out_path = (output_dir / file_id).as_posix()
command = [
"ffmpeg",
"-f",
@@ -240,28 +233,12 @@ class MTB_ExportWithFfmpeg:
process.stdin.close()
process.wait()
return (out_path,)
else:
if has_alpha:
if codec in ["prores_ks", "libx264", "libx265"]:
pix_fmt = (
"yuva444p" if codec == "prores_ks" else "yuva420p"
)
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]
frames = [frame.astype(np.uint16) * 257 for frame in frames]
height, width, _ = frames[0].shape
out_path = (output_dir / file_id).as_posix()
# Prepare the FFmpeg command
command = [
@@ -281,26 +258,17 @@ class MTB_ExportWithFfmpeg:
"-",
"-c:v",
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)
pbar = comfy.utils.ProgressBar(len(frames))
for frame in frames:
model_management.throw_exception_if_processing_interrupted()
process.stdin.write(frame.tobytes())
pbar.update(1)
process.stdin.close()
process.wait()
@@ -312,9 +280,9 @@ def prepare_animated_batch(
batch: torch.Tensor,
pingpong=False,
resize_by=1.0,
resample_filter: Image.Resampling | None = None,
resample_filter: Optional[Image.Resampling] = None,
image_type=np.uint8,
) -> list[Image.Image]:
) -> List[Image.Image]:
images = tensor2np(batch)
images = [frame.astype(image_type) for frame in images]
@@ -340,10 +308,7 @@ def prepare_animated_batch(
# todo: deprecate for apng
class MTB_SaveGif:
"""Save the images from the batch as a GIF.
[DEPRACATED] Use VHS nodes instead
"""
"""Save the images from the batch as a GIF"""
@classmethod
def INPUT_TYPES(cls):
@@ -363,7 +328,6 @@ class MTB_SaveGif:
OUTPUT_NODE = True
CATEGORY = "mtb/IO"
FUNCTION = "save_gif"
DEPRECATED = True
def save_gif(
self,
-161
View File
@@ -1,161 +0,0 @@
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
]
+1 -2
View File
@@ -1,5 +1,6 @@
import comfy.utils
from PIL import Image
from rembg import remove
from ..utils import pil2tensor, tensor2pil
@@ -63,8 +64,6 @@ class MTB_ImageRemoveBackgroundRembg:
post_process_mask,
bgcolor,
):
from rembg import remove
pbar = comfy.utils.ProgressBar(image.size(0))
images = tensor2pil(image)
-351
View File
@@ -1,351 +0,0 @@
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]
+1 -30
View File
@@ -45,19 +45,6 @@ class MTB_TransformImage:
),
"constant_color": ("COLOR", {"default": "#000000"}),
},
"optional": {
"filter_type": (
[
"nearest",
"box",
"bilinear",
"hamming",
"bicubic",
"lanczos",
],
{"default": "bilinear"},
),
},
}
FUNCTION = "transform"
@@ -74,18 +61,7 @@ class MTB_TransformImage:
shear: float,
border_handling="edge",
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)
y = int(y)
angle = int(angle)
@@ -139,12 +115,7 @@ class MTB_TransformImage:
img = cast(
Image.Image,
TF.affine(
img,
angle=angle,
scale=zoom,
translate=[x, y],
shear=shear,
interpolation=resampling_filter,
img, angle=angle, scale=zoom, translate=[x, y], shear=shear
),
)
+179 -181
View File
@@ -1,181 +1,179 @@
[build-system]
requires = ["setuptools", "wheel"]
build-backend = "setuptools.build_meta"
[project]
name = "comfy-mtb"
version = "0.2.1"
description = "Animation oriented nodes pack for ComfyUI."
license = "MIT"
readme = "README.md"
# repository = ""
# url = "https://github.com/melMass/comfy_mtb"
authors = [{ name = "Mel Massadian", email = "mel@melmassadian.com" }]
classifiers = [
"License :: OSI Approved :: MIT License",
"Operating System :: OS Independent",
"Programming Language :: Python",
"Programming Language :: Python :: 3",
"Programming Language :: Python :: 3.10",
"Programming Language :: Python :: 3.11",
"Intended Audience :: Developers",
]
requires-python = ">=3.10"
dependencies = [
"qrcode",
"cachetools",
"onnxruntime-gpu",
"requirements-parserx",
"rembg",
"imageio_ffmpeg",
"rich",
"rich_argparse",
"matplotlib",
"pillow",
]
optional-dependencies = { mel = [
"jupyterlab==4.1.6",
], dev = [
"black[jupyter]",
"codespell",
"marimo",
"mypy",
"pre-commit",
"pytest",
"pytest-cov",
"pytest-random-order",
"ruff",
], doc = [
"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"
Repository = "https://github.com/melMass/comfy_mtb"
Issues = "https://github.com/melMass/comfy_mtb/issues"
[tool.comfy]
PublisherId = "mel"
DisplayName = "comfy-mtb"
Icon = "https://avatars.githubusercontent.com/u/7041726?v=4"
[tool.bumpversion]
current_version = "0.2.1"
parse = "(?P<major>\\d+)\\.(?P<minor>\\d+)\\.(?P<patch>\\d+)"
serialize = ["{major}.{minor}.{patch}"]
search = "{current_version}"
replace = "{new_version}"
regex = false
ignore_missing_version = false
ignore_missing_files = false
tag = true
sign_tags = true
tag_name = "v{new_version}"
tag_message = "⬆️ Bump version: {current_version} → {new_version}"
allow_dirty = true
commit = true
message = "⬆️ Bump version: {current_version} → {new_version}"
commit_args = ""
[[tool.bumpversion.files]]
filename = "__init__.py"
search = "__version__ = \"{current_version}\""
replace = "__version__ = \"{new_version}\""
[[tool.bumpversion.files]]
filename = "pyproject.toml"
search = "version = \"{current_version}\""
replace = "version = \"{new_version}\""
# [[tool.bumpversion.files]]
# filename = "your_package/__init__.py"
# search = "__version__ = '{current_version}'"
# replace = "__version__ = '{new_version}'"
# INFO: All those remaining keys are meant for local dev
[tool.pyright]
include = ["."]
exclude = [
"**/node_modules",
"**/__pycache__",
"src/experimental",
"src/typestubs",
]
ignore = ["src/oldstuff"]
defineConstant = { DEBUG = true }
extraPaths = ["python", "../.."]
stubPath = "src/stubs"
reportMissingImports = true
reportMissingTypeStubs = false
typeCheckingMode = "basic"
pythonVersion = "3.10"
pythonPlatform = "Windows"
[tool.pytest.ini_options]
log_level = "DEBUG"
log_cli = true
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]
profile = "black"
line_length = 88
auto_identify_namespace_packages = false
# NOTE:
# 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
force_single_line = false
known_first_party = ["mtb"]
extend_skip = ["archives"]
combine_straight_imports = true
[tool.coverage.run]
parallel = true
source = ["docs", "tests", "comfy-mtb"]
[tool.coverage.report]
fail_under = 90
show_missing = true
[tool.coverage.html]
show_contexts = true
[tool.ruff]
line-length = 79
select = ["A", "B", "C", "D", "E", "F", "FBT", "I", "N", "S", "SIM", "UP", "W"]
# NOTE:
# D102 - undocumented-public-method (noisy)
# D103 - undocumented-public-function (noisy)
# D100 - undocumented-public-module (noisy)
# N802 - invalid-function-name (forced by comfy's arch)
ignore = ["D103", "D102", "D100", "N802"]
# exclude auto generated file
extend-exclude = ["./docs/conf.py"]
[tool.ruff.per-file-ignores]
# imported but unused
"__init__.py" = ["F401"]
# use of assert detected
"tests/*" = ["S101"]
[tool.ruff.pydocstyle]
convention = "numpy"
[tool.mypy]
pretty = true
ignore_missing_imports = true
# exclude auto generated file
exclude = ["docs/conf.py"]
[tool.codespell]
# exclude auto generated file
skip = "./docs/conf.py,poetry.lock"
check-filenames = true
[build-system]
requires = ["setuptools", "wheel"]
build-backend = "setuptools.build_meta"
[project]
name = "comfy-mtb"
version = "0.1.6"
description = "Animation oriented nodes pack for ComfyUI."
license = "MIT"
readme = "README.md"
# repository = ""
# url = "https://github.com/melMass/comfy_mtb"
authors = [{ name = "Mel Massadian", email = "mel@melmassadian.com" }]
classifiers = [
"License :: OSI Approved :: MIT License",
"Operating System :: OS Independent",
"Programming Language :: Python",
"Programming Language :: Python :: 3",
"Programming Language :: Python :: 3.10",
"Programming Language :: Python :: 3.11",
"Intended Audience :: Developers",
]
requires-python = ">=3.10"
dependencies = [
"qrcode",
"onnxruntime-gpu",
"requirements-parserx",
"rembg",
"imageio_ffmpeg",
"rich",
"rich_argparse",
"matplotlib",
"pillow",
]
optional-dependencies = { mel = [
"jupyterlab==4.1.6",
], dev = [
"black[jupyter]",
"codespell",
"mypy",
"pre-commit",
"pytest",
"pytest-cov",
"pytest-random-order",
"ruff",
], doc = [
"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"
Repository = "https://github.com/melMass/comfy_mtb"
Issues = "https://github.com/melMass/comfy_mtb/issues"
[tool.comfy]
PublisherId = "mel"
DisplayName = "comfy-mtb"
Icon = "https://avatars.githubusercontent.com/u/7041726?v=4"
[tool.bumpversion]
current_version = "0.1.6"
parse = "(?P<major>\\d+)\\.(?P<minor>\\d+)\\.(?P<patch>\\d+)"
serialize = ["{major}.{minor}.{patch}"]
search = "{current_version}"
replace = "{new_version}"
regex = false
ignore_missing_version = false
ignore_missing_files = false
tag = true
sign_tags = true
tag_name = "v{new_version}"
tag_message = "⬆️ Bump version: {current_version} → {new_version}"
allow_dirty = true
commit = true
message = "⬆️ Bump version: {current_version} → {new_version}"
commit_args = ""
[[tool.bumpversion.files]]
filename = "__init__.py"
search = "__version__ = \"{current_version}\""
replace = "__version__ = \"{new_version}\""
[[tool.bumpversion.files]]
filename = "pyproject.toml"
search = "version = \"{current_version}\""
replace = "version = \"{new_version}\""
# [[tool.bumpversion.files]]
# filename = "your_package/__init__.py"
# search = "__version__ = '{current_version}'"
# replace = "__version__ = '{new_version}'"
# INFO: All those remaining keys are meant for local dev
[tool.pyright]
include = ["."]
exclude = [
"**/node_modules",
"**/__pycache__",
"src/experimental",
"src/typestubs",
]
ignore = ["src/oldstuff"]
defineConstant = { DEBUG = true }
extraPaths = ["python", "../.."]
stubPath = "src/stubs"
reportMissingImports = true
reportMissingTypeStubs = false
typeCheckingMode = "basic"
pythonVersion = "3.10"
pythonPlatform = "Windows"
[tool.pytest.ini_options]
log_level = "DEBUG"
log_cli = true
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]
profile = "black"
line_length = 88
auto_identify_namespace_packages = false
# NOTE:
# 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
force_single_line = false
known_first_party = ["mtb"]
extend_skip = ["archives"]
combine_straight_imports = true
[tool.coverage.run]
parallel = true
source = ["docs", "tests", "comfy-mtb"]
[tool.coverage.report]
fail_under = 90
show_missing = true
[tool.coverage.html]
show_contexts = true
[tool.ruff]
line-length = 79
select = ["A", "B", "C", "D", "E", "F", "FBT", "I", "N", "S", "SIM", "UP", "W"]
# NOTE:
# D102 - undocumented-public-method (noisy)
# D103 - undocumented-public-function (noisy)
# D100 - undocumented-public-module (noisy)
# N802 - invalid-function-name (forced by comfy's arch)
ignore = ["D103", "D102", "D100", "N802"]
# exclude auto generated file
extend-exclude = ["./docs/conf.py"]
[tool.ruff.per-file-ignores]
# imported but unused
"__init__.py" = ["F401"]
# use of assert detected
"tests/*" = ["S101"]
[tool.ruff.pydocstyle]
convention = "numpy"
[tool.mypy]
pretty = true
ignore_missing_imports = true
# exclude auto generated file
exclude = ["docs/conf.py"]
[tool.codespell]
# exclude auto generated file
skip = "./docs/conf.py,poetry.lock"
check-filenames = true
-1
View File
@@ -8,4 +8,3 @@ rich
rich_argparse
matplotlib
pillow
cachetools
+9 -78
View File
@@ -2,7 +2,6 @@ import contextlib
import functools
import importlib
import math
import operator
import os
import shlex
import shutil
@@ -12,7 +11,6 @@ import sys
import uuid
from collections.abc import Callable, Sequence
from enum import Enum
from functools import reduce
from pathlib import Path
from typing import TypeVar
@@ -165,9 +163,9 @@ class IPChecker:
def __init__(self):
self.ips = list(self.get_local_ips())
log.debug(f"Found {len(self.ips)} local ips")
self.checked_ips: set[str] = set()
self.checked_ips = set()
def get_working_ip(self, test_url_template: str):
def get_working_ip(self, test_url_template):
for ip in self.ips:
if ip not in self.checked_ips:
self.checked_ips.add(ip)
@@ -177,7 +175,7 @@ class IPChecker:
return None
@staticmethod
def get_local_ips(prefix: str = "192.168."):
def get_local_ips(prefix="192.168."):
hostname = socket.gethostname()
log.debug(f"Getting local ips for {hostname}")
for info in socket.getaddrinfo(hostname, None):
@@ -187,9 +185,9 @@ class IPChecker:
if info[0] == socket.AF_INET and info[4][0].startswith(prefix):
yield info[4][0]
def _test_url(self, url: str):
def _test_url(self, url):
try:
response = requests.get(url, timeout=10)
response = requests.get(url)
return response.status_code == 200
except Exception:
return False
@@ -200,7 +198,7 @@ def get_server_info():
from comfy.cli_args import args
ip_checker = IPChecker()
base_url: str = args.listen
base_url = args.listen
if base_url == "0.0.0.0":
log.debug("Server set to 0.0.0.0, we will try to resolve the host IP")
base_url = ip_checker.get_working_ip(
@@ -214,37 +212,6 @@ def get_server_info():
# 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
@@ -498,12 +465,8 @@ here = Path(__file__).parent.absolute()
# - Construct the absolute path to the ComfyUI directory
comfy_dir = Path(folder_paths.base_path)
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)
input_dir = Path(folder_paths.input_directory)
styles_dir = comfy_dir / "styles"
session_id = str(uuid.uuid4())
# - Construct the path to the font file
@@ -513,10 +476,9 @@ font_path = here / "data" / "font.ttf"
extern_root = here / "extern"
add_path(extern_root)
if extern_root.exists():
for pth in extern_root.iterdir():
if pth.is_dir():
add_path(pth)
for pth in extern_root.iterdir():
if pth.is_dir():
add_path(pth)
# - Add the ComfyUI directory and custom nodes path to the sys.path list
add_path(comfy_dir)
@@ -630,37 +592,6 @@ def tensor2np(tensor: torch.Tensor) -> list[npt.NDArray[np.uint8]]:
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):
pad_width = np.array(((0, 0), (top, bottom), (left, right)))
print(
+54 -92
View File
@@ -1,9 +1,10 @@
/**
* @module Shared utilities
* File: comfy_shared.js
* Project: comfy_mtb
* Author: Mel Massadian
*
* Copyright (c) 2023-2024 Mel Massadian
*
*/
// Reference the shared typedefs file
@@ -260,16 +261,6 @@ 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 {LLink} link
@@ -371,37 +362,23 @@ export function getWidgetType(config) {
* @param {NodeType} nodeType The nodetype to attach the documentation to
* @param {str} prefix A prefix added to each dynamic inputs
* @param {str | [str]} inputType The datatype(s) of those dynamic inputs
* @param {{separator?:string, start_index?:number, link?:LLink, ioSlot?:INodeInputSlot | INodeOutputSlot}?} [opts] Extra options
* @param {{link?:LLink, ioSlot?:INodeInputSlot | INodeOutputSlot}?} opts
* @returns
*/
export const setupDynamicConnections = (
nodeType,
prefix,
inputType,
opts = undefined,
) => {
export const setupDynamicConnections = (nodeType, prefix, inputType, opts) => {
infoLogger(
'Setting up dynamic connections for',
Object.getOwnPropertyDescriptors(nodeType).title.value,
)
/** @type {{separator:string, start_index:number, link?:LLink, ioSlot?:INodeInputSlot | INodeOutputSlot}?} */
const options = Object.assign(
{
separator: '_',
start_index: 1,
},
opts || {},
)
/** @type {{link?:LLink, ioSlot?:INodeInputSlot | INodeOutputSlot}} */
const options = opts || {}
const onNodeCreated = nodeType.prototype.onNodeCreated
const inputList = typeof inputType === 'object'
nodeType.prototype.onNodeCreated = function () {
const r = onNodeCreated ? onNodeCreated.apply(this, []) : undefined
this.addInput(
`${prefix}${options.separator}${options.start_index}`,
inputList ? '*' : inputType,
)
this.addInput(`${prefix}_1`, inputList ? '*' : inputType)
return r
}
@@ -436,7 +413,7 @@ export const setupDynamicConnections = (
this,
slotIndex,
isConnected,
`${prefix}${options.separator}`,
`${prefix}_`,
inputType,
options,
)
@@ -452,7 +429,7 @@ export const setupDynamicConnections = (
* @param {bool} connected - Was this event connecting or disconnecting
* @param {string} [connectionPrefix] - The common prefix of the dynamic inputs
* @param {string|[string]} [connectionType] - The type of the dynamic connection
* @param {{start_index?:number, link?:LLink, ioSlot?:INodeInputSlot | INodeOutputSlot}} [opts] - extra options
* @param {{link?:LLink, ioSlot?:INodeInputSlot | INodeOutputSlot}} [opts] - extra options
*/
export const dynamic_connection = (
node,
@@ -462,18 +439,13 @@ export const dynamic_connection = (
connectionType = '*',
opts = undefined,
) => {
/* {{start_index:number, link?:LLink, ioSlot?:INodeInputSlot | INodeOutputSlot}} [opts] - extra options*/
const options = Object.assign(
{
start_index: 1,
},
opts || {},
)
/* @type {{link?:LLink, ioSlot?:INodeInputSlot | INodeOutputSlot}} [opts] - extra options*/
const options = opts || {}
// function to test if input is a dynamic one
const isDynamicInput = (inputName) => inputName.startsWith(connectionPrefix)
if (node.inputs.length > 0 && !isDynamicInput(node.inputs[index].name)) {
if (
node.inputs.length > 0 &&
!node.inputs[index].name.startsWith(connectionPrefix)
) {
return
}
@@ -492,7 +464,7 @@ export const dynamic_connection = (
const to_remove = []
for (let n = 1; n < node.inputs.length; n++) {
const element = node.inputs[n]
if (!element.link && isDynamicInput(element.name)) {
if (!element.link) {
if (node.widgets) {
const w = node.widgets.find((w) => w.name === element.name)
if (w) {
@@ -518,25 +490,14 @@ export const dynamic_connection = (
infoLogger('Cleaning inputs: making it sequential again')
// make inputs sequential again
let prefixed_idx = options.start_index
for (let i = 0; i < node.inputs.length; i++) {
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
}
let name = `${connectionPrefix}${i + 1}`
if (nameArray.length > 0) {
name = i < nameArray.length ? nameArray[i] : name
}
// preserve label if it exists
node.inputs[i].label = node.inputs[i].label || name
node.inputs[i].label = name
node.inputs[i].name = name
}
}
@@ -576,16 +537,11 @@ export const dynamic_connection = (
if (node.inputs.length === 0) return
// add an extra input
if (node.inputs[node.inputs.length - 1].link !== null) {
// count only the prefixed inputs
const nextIndex = node.inputs.reduce(
(acc, cur) => (isDynamicInput(cur.name) ? ++acc : acc),
0,
)
const nextIndex = node.inputs.length
const name =
nextIndex < nameArray.length
? nameArray[nextIndex]
: `${connectionPrefix}${nextIndex + options.start_index}`
: `${connectionPrefix}${nextIndex + 1}`
infoLogger(`Adding input ${nextIndex + 1} (${name})`)
node.addInput(name, conType)
@@ -675,6 +631,39 @@ 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
// #region documentation widget
@@ -1129,33 +1118,7 @@ export const addDeprecation = (nodeType, reason) => {
// #endregion
// #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
// #region API / graph utilities
export const getAPIInputs = () => {
const inputs = {}
let counter = 1
@@ -1208,4 +1171,3 @@ export const getNodes = (skip_unused) => {
}
return nodes
}
// #endregion
+295 -296
View File
@@ -13,40 +13,40 @@ import { api } from '../../scripts/api.js'
import { app } from '../../scripts/app.js'
import { LocalStorageManager } from './comfy_shared.js'
const styles = {
lighbox: {
position: 'fixed',
top: 0,
left: 0,
width: '100vw',
height: '100vh',
background: 'rgba(0,0,0,0.5)',
display: 'none',
justifyContent: 'center',
alignItems: 'center',
zIndex: 999,
},
lightboxBtn: (extra) => ({
position: 'absolute',
top: '50%',
background: 'none',
border: 'none',
color: '#fff',
zIndex: 1000,
fontSize: '30px',
cursor: 'pointer',
pointerEvents: 'auto',
...extra,
}),
img_list: {
minHeight: '30px',
maxHeight: '300px',
width: '100vw',
position: 'absolute',
bottom: 0,
zIndex: 10,
background: '#333',
overflow: 'auto',
},
lighbox: {
position: 'fixed',
top: 0,
left: 0,
width: '100vw',
height: '100vh',
background: 'rgba(0,0,0,0.5)',
display: 'none',
justifyContent: 'center',
alignItems: 'center',
zIndex: 999,
},
lightboxBtn: (extra) => ({
position: 'absolute',
top: '50%',
background: 'none',
border: 'none',
color: '#fff',
zIndex: 1000,
fontSize: '30px',
cursor: 'pointer',
pointerEvents: 'auto',
...extra,
}),
img_list: {
minHeight: '30px',
maxHeight: '300px',
width: '100vw',
position: 'absolute',
bottom: 0,
zIndex: 10,
background: '#333',
overflow: 'auto',
},
}
let currentImageIndex = 0
@@ -58,299 +58,298 @@ const storage = new LocalStorageManager('mtb')
let activated = storage.get('image_feed', false)
app.registerExtension({
name: 'mtb.ImageFeed',
setup: () => {
app.ui.settings.addSetting({
id: 'mtb.Main.image-feed-enabled',
category: ['mtb', 'Main', 'image-feed-enabled'],
name: 'Enable Image Feed',
type: 'boolean',
defaultValue: false,
attrs: {
style: {
fontFamily: 'monospace',
},
},
async onChange(value) {
storage.set('image_feed', value)
activated = value
},
})
},
init: async () => {
if (!activated) {
return
}
const pythongossFeed = app.extensions.find(
(e) => e.name === 'pysssss.ImageFeed',
)
if (pythongossFeed) {
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
}
// - HTML & CSS
//- lightbox
const lightboxContainer = document.createElement('div')
Object.assign(lightboxContainer.style, styles.lighbox)
name: 'mtb.ImageFeed',
setup: () => {
app.ui.settings.addSetting({
id: 'mtb.imageFeed.enabled',
name: '[⚡mtb] Enable image feed',
type: 'boolean',
defaultValue: true,
attrs: {
style: {
fontFamily: 'monospace',
},
},
async onChange(value) {
storage.set('image_feed', value)
activated = value
},
})
},
init: async () => {
if (!activated) {
return
}
const pythongossFeed = app.extensions.find(
(e) => e.name === 'pysssss.ImageFeed',
)
if (pythongossFeed) {
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
}
// - HTML & CSS
//- lightbox
const lightboxContainer = document.createElement('div')
Object.assign(lightboxContainer.style, styles.lighbox)
const lightboxImage = document.createElement('img')
Object.assign(lightboxImage.style, {
maxHeight: '100%',
maxWidth: '100%',
borderRadius: '5px',
})
const lightboxImage = document.createElement('img')
Object.assign(lightboxImage.style, {
maxHeight: '100%',
maxWidth: '100%',
borderRadius: '5px',
})
// previous and next buttons
const lightboxPrevBtn = document.createElement('button')
const lightboxNextBtn = document.createElement('button')
// previous and next buttons
const lightboxPrevBtn = document.createElement('button')
const lightboxNextBtn = document.createElement('button')
lightboxPrevBtn.textContent = '❮'
lightboxNextBtn.textContent = '❯'
lightboxPrevBtn.textContent = '❮'
lightboxNextBtn.textContent = '❯'
Object.assign(lightboxPrevBtn.style, styles.lightboxBtn({ left: '0%' }))
Object.assign(lightboxNextBtn.style, styles.lightboxBtn({ right: '0%' }))
Object.assign(lightboxPrevBtn.style, styles.lightboxBtn({ left: '0%' }))
Object.assign(lightboxNextBtn.style, styles.lightboxBtn({ right: '0%' }))
// close button
const lightboxCloseBtn = document.createElement('button')
Object.assign(
lightboxCloseBtn.style,
styles.lightboxBtn({ right: '0', top: '0' }),
)
lightboxCloseBtn.textContent = '❌'
// close button
const lightboxCloseBtn = document.createElement('button')
Object.assign(
lightboxCloseBtn.style,
styles.lightboxBtn({ right: '0', top: '0' }),
)
lightboxCloseBtn.textContent = '❌'
const lightboxButtons = document.createElement('div')
Object.assign(lightboxButtons.style, {
position: 'absolute',
top: '0%',
right: '0%',
// transform: "translate(50%, -50%)",
height: '100%',
width: '100%',
background: 'none',
border: 'none',
color: '#fff',
fontSize: '30px',
cursor: 'pointer',
pointerEvents: 'none',
})
const lightboxButtons = document.createElement('div')
Object.assign(lightboxButtons.style, {
position: 'absolute',
top: '0%',
right: '0%',
// transform: "translate(50%, -50%)",
height: '100%',
width: '100%',
background: 'none',
border: 'none',
color: '#fff',
fontSize: '30px',
cursor: 'pointer',
pointerEvents: 'none',
})
lightboxButtons.append(lightboxPrevBtn, lightboxNextBtn, lightboxCloseBtn)
lightboxContainer.append(lightboxButtons, lightboxImage)
lightboxButtons.append(lightboxPrevBtn, lightboxNextBtn, lightboxCloseBtn)
lightboxContainer.append(lightboxButtons, lightboxImage)
//- image list
const imageListContainer = document.createElement('div')
Object.assign(imageListContainer.style, styles.img_list)
//- image list
const imageListContainer = document.createElement('div')
Object.assign(imageListContainer.style, styles.img_list)
const createImgListBtn = (text, style) => {
const btn = document.createElement('button')
btn.type = 'button'
btn.textContent = text
Object.assign(btn.style, {
...style,
border: 'none',
color: '#fff',
background: 'none',
height: '20px',
cursor: 'pointer',
position: 'absolute',
top: '5px',
fontSize: '12px',
lineHeight: '12px',
})
imageListContainer.append(btn)
return btn
}
const showBtn = document.createElement('button')
const closeBtn = createImgListBtn('❌', {
width: '20px',
textIndent: '-4px',
right: '5px',
})
const loadButton = createImgListBtn('Load Session History', {
right: '90px',
})
const clearButton = createImgListBtn('Clear', {
right: '30px',
})
const createImgListBtn = (text, style) => {
const btn = document.createElement('button')
btn.type = 'button'
btn.textContent = text
Object.assign(btn.style, {
...style,
border: 'none',
color: '#fff',
background: 'none',
height: '20px',
cursor: 'pointer',
position: 'absolute',
top: '5px',
fontSize: '12px',
lineHeight: '12px',
})
imageListContainer.append(btn)
return btn
}
const showBtn = document.createElement('button')
const closeBtn = createImgListBtn('❌', {
width: '20px',
textIndent: '-4px',
right: '5px',
})
const loadButton = createImgListBtn('Load Session History', {
right: '90px',
})
const clearButton = createImgListBtn('Clear', {
right: '30px',
})
//- tools popup button
showBtn.classList.add('comfy-settings-btn')
Object.assign(showBtn.style, {
right: '16px',
cursor: 'pointer',
display: 'none',
})
//- tools popup button
showBtn.classList.add('comfy-settings-btn')
Object.assign(showBtn.style, {
right: '16px',
cursor: 'pointer',
display: 'none',
})
//- append to DOM
document.body.append(imageListContainer)
//- append to DOM
document.body.append(imageListContainer)
showBtn.textContent = '🖼'
showBtn.onclick = () => {
imageListContainer.style.display = 'block'
showBtn.style.display = 'none'
}
document.querySelector('.comfy-settings-btn').after(showBtn)
document.querySelector('.comfy-settings-btn').after(lightboxContainer)
showBtn.textContent = '🖼'
showBtn.onclick = () => {
imageListContainer.style.display = 'block'
showBtn.style.display = 'none'
}
document.querySelector('.comfy-settings-btn').after(showBtn)
document.querySelector('.comfy-settings-btn').after(lightboxContainer)
// for (const { output } of history) {
// if (output?.images) {
// for (const src of output.images) {
// const img = document.createElement("img");
// const but = document.createElement("button");
// for (const { output } of history) {
// if (output?.images) {
// for (const src of output.images) {
// const img = document.createElement("img");
// const but = document.createElement("button");
//- callbacks
closeBtn.onclick = () => {
imageListContainer.style.display = 'none'
showBtn.style.display = 'unset'
}
//- callbacks
closeBtn.onclick = () => {
imageListContainer.style.display = 'none'
showBtn.style.display = 'unset'
}
clearButton.onclick = () => {
imageListContainer.replaceChildren(closeBtn, clearButton, loadButton)
}
clearButton.onclick = () => {
imageListContainer.replaceChildren(closeBtn, clearButton, loadButton)
}
lightboxNextBtn.onclick = () => {
currentImageIndex = (currentImageIndex + 1) % imageUrls.length
const imageUrl = imageUrls[currentImageIndex]
lightboxImage.src = imageUrl
}
lightboxNextBtn.onclick = () => {
currentImageIndex = (currentImageIndex + 1) % imageUrls.length
const imageUrl = imageUrls[currentImageIndex]
lightboxImage.src = imageUrl
}
// Modify the lightboxPrevBtn onclick callback
lightboxPrevBtn.onclick = () => {
currentImageIndex =
(currentImageIndex - 1 + imageUrls.length) % imageUrls.length
const imageUrl = imageUrls[currentImageIndex]
lightboxImage.src = imageUrl
}
// Modify the lightboxPrevBtn onclick callback
lightboxPrevBtn.onclick = () => {
currentImageIndex =
(currentImageIndex - 1 + imageUrls.length) % imageUrls.length
const imageUrl = imageUrls[currentImageIndex]
lightboxImage.src = imageUrl
}
lightboxCloseBtn.onclick = () => {
lightboxContainer.style.display = 'none'
}
lightboxImage.onclick = lightboxNextBtn.onclick
/**
* 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
* the image in the lightbox.
* @param {*} src
*/
const createImageBtn = (src) => {
console.debug(`making image ${src.filename}`)
const img = document.createElement('img')
const but = document.createElement('button')
lightboxCloseBtn.onclick = () => {
lightboxContainer.style.display = 'none'
}
lightboxImage.onclick = lightboxNextBtn.onclick
/**
* 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
* the image in the lightbox.
* @param {*} src
*/
const createImageBtn = (src) => {
console.debug(`making image ${src.filename}`)
const img = document.createElement('img')
const but = document.createElement('button')
Object.assign(but.style, {
height: '120px',
width: '120px',
border: 'none',
padding: 0,
margin: 0,
})
Object.assign(img.style, {
width: '100%',
height: '100%',
objectFit: 'cover',
})
Object.assign(but.style, {
height: '120px',
width: '120px',
border: 'none',
padding: 0,
margin: 0,
})
Object.assign(img.style, {
width: '100%',
height: '100%',
objectFit: 'cover',
})
img.src = `/view?filename=${encodeURIComponent(src.filename)}&type=${
src.type
}&subfolder=${encodeURIComponent(src.subfolder)}`
img.src = `/view?filename=${encodeURIComponent(src.filename)}&type=${
src.type
}&subfolder=${encodeURIComponent(src.subfolder)}`
imageUrls.push(img.src)
imageUrls.push(img.src)
console.debug(img.src)
console.debug(img.src)
img.onload = () => {
but.style.width = `${120 * (img.naturalWidth / img.naturalHeight)}px`
}
img.onload = () => {
but.style.width = `${120 * (img.naturalWidth / img.naturalHeight)}px`
}
but.onclick = () => {
lightboxContainer.style.display = 'flex'
// add the same image to the lightbox
lightboxImage.src = img.src
// lighboxContainer.replaceChildren(lightboxButtons, img);
}
but.onclick = () => {
lightboxContainer.style.display = 'flex'
// add the same image to the lightbox
lightboxImage.src = img.src
// lighboxContainer.replaceChildren(lightboxButtons, img);
}
// add right click menu
but.addEventListener('contextmenu', (e) => {
e.preventDefault()
// add right click menu
but.addEventListener('contextmenu', (e) => {
e.preventDefault()
if (image_menu) {
image_menu.remove()
}
if (image_menu) {
image_menu.remove()
}
image_menu = document.createElement('div')
Object.assign(image_menu.style, {
position: 'absolute',
top: `${e.clientY}px`,
left: `${e.clientX}px`,
background: '#333',
color: '#fff',
padding: '5px',
borderRadius: '5px',
zIndex: 999,
})
const load_img = document.createElement('button')
load_img.textContent = 'Load'
load_img.onclick = () => {
app.handleFile(img.src)
}
image_menu = document.createElement('div')
Object.assign(image_menu.style, {
position: 'absolute',
top: `${e.clientY}px`,
left: `${e.clientX}px`,
background: '#333',
color: '#fff',
padding: '5px',
borderRadius: '5px',
zIndex: 999,
})
const load_img = document.createElement('button')
load_img.textContent = 'Load'
load_img.onclick = () => {
app.handleFile(img.src)
}
image_menu.appendChild(load_img)
document.body.appendChild(image_menu)
})
image_menu.appendChild(load_img)
document.body.appendChild(image_menu)
})
but.append(img)
imageListContainer.prepend(but)
}
but.append(img)
imageListContainer.prepend(but)
}
loadButton.onclick = async () => {
const all_history = await api.getHistory()
for (const history of all_history.History) {
if (history.outputs) {
for (const key of Object.keys(history.outputs)) {
console.debug(key)
if (history.outputs[key].images) {
for (const im of history.outputs[key].images) {
console.debug(im)
createImageBtn(im)
}
}
}
// for (const src of outputs.outputs.images) {
// console.debug(src)
// makeImage(`${src.subfolder}/${src.filename}`)
// }
}
}
}
loadButton.onclick = async () => {
const all_history = await api.getHistory()
for (const history of all_history.History) {
if (history.outputs) {
for (const key of Object.keys(history.outputs)) {
console.debug(key)
if (history.outputs[key].images) {
for (const im of history.outputs[key].images) {
console.debug(im)
createImageBtn(im)
}
}
}
// for (const src of outputs.outputs.images) {
// console.debug(src)
// makeImage(`${src.subfolder}/${src.filename}`)
// }
}
}
}
///////-------
///////-------
// const all_history = await api.getHistory()
// for (const history of all_history.History) {
// if (history.outputs) {
// for (const key of Object.keys(history.outputs)) {
// for (const im of history.outputs[key].images) {
// makeImage(im)
// }
// }
// // for (const src of outputs.outputs.images) {
// // console.debug(src)
// // makeImage(`${src.subfolder}/${src.filename}`)
// // }
// }
// }
// const all_history = await api.getHistory()
// for (const history of all_history.History) {
// if (history.outputs) {
// for (const key of Object.keys(history.outputs)) {
// for (const im of history.outputs[key].images) {
// makeImage(im)
// }
// }
// // for (const src of outputs.outputs.images) {
// // console.debug(src)
// // makeImage(`${src.subfolder}/${src.filename}`)
// // }
// }
// }
//- Hook into the API
api.addEventListener('executed', ({ detail }) => {
if (detail?.output?.images) {
for (const src of detail.output.images) {
console.debug(`Adding ${src} to image feed`)
createImageBtn(src)
}
}
})
},
//- Hook into the API
api.addEventListener('executed', ({ detail }) => {
if (detail?.output?.images) {
for (const src of detail.output.images) {
console.debug(`Adding ${src} to image feed`)
createImageBtn(src)
}
}
})
},
})
-339
View File
@@ -1,339 +0,0 @@
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
@@ -1,28 +0,0 @@
// 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
@@ -1,504 +0,0 @@
/**
* 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
}
+1095 -1155
View File
File diff suppressed because it is too large Load Diff
+1 -1
Submodule wiki updated: fa7fec28a3...4db733ae92