Compare commits

..
Author SHA1 Message Date
shanexi 6a27517db6 Convert to ShellAgent with default value assigned 2024-10-21 17:33:59 +08:00
wl-zhao 68e5dad893 update file path 2024-10-20 11:23:02 +08:00
wl-zhao 9a606f550d update file_upload 2024-10-20 10:43:40 +08:00
wl-zhao 1b5fd5e7d5 update file_upload 2024-10-20 10:41:12 +08:00
wl-zhao 44ac4cd0f8 update relpath of models 2024-10-20 10:19:22 +08:00
wl-zhao 83e1caa5fa update relpath of models 2024-10-20 10:04:06 +08:00
wl-zhao 58200be879 fix folder_path bugs 2024-10-20 09:37:32 +08:00
wl-zhao 895fc30c37 fix folder_path bugs 2024-10-20 09:33:56 +08:00
wl-zhao 62ab265a80 update catch git error 2024-10-20 09:21:54 +08:00
wl-zhao f1bf4485c1 update validate inputs 2024-10-18 14:05:58 +08:00
wl-zhao b8292f45ba update commit of node deps 2024-10-18 13:54:52 +08:00
wl-zhao 166fb12aba update node_deps_info 2024-10-18 13:37:07 +08:00
wl-zhao 3f6f36cd81 update node_deps_info 2024-10-18 12:25:48 +08:00
wl-zhao 0982abfcf8 update node_deps_info 2024-10-18 12:18:52 +08:00
wl-zhao e06fb3d5d3 update model suffix 2024-10-18 10:25:54 +08:00
wl-zhao 6afc9ce442 use glob to traverse 2024-10-18 10:16:30 +08:00
wl-zhao 0d9bd35d84 Merge branch 'main' of github.com:myshell-ai/ComfyUI-ShellAgent-Plugin into main 2024-10-18 10:05:45 +08:00
wl-zhao f49589ee82 add printed logs 2024-10-18 10:05:36 +08:00
tiancheng c33e93af11 README and license 2024-10-18 01:58:27 +03:00
wl-zhao 60d68d99ef fix input image bug 2024-10-17 23:40:42 +08:00
wl-zhao b8b6657801 fix input image bug 2024-10-17 23:38:46 +08:00
wl-zhao ad8715f3a2 fix input image bug 2024-10-17 23:21:31 +08:00
wl-zhao 92ab0995e4 fix input image bug 2024-10-17 23:12:36 +08:00
Wenliang Zhao e146e2512f Merge pull request #2 from myshell-ai/1-optimize-custom-node
1 optimize custom node
2024-10-17 16:36:50 +08:00
wl-zhao 49fb34f250 update error message 2024-10-17 16:26:45 +08:00
wl-zhao f2926099fe update error message 2024-10-17 16:22:49 +08:00
shanexi fb4e608fc6 Fix bugs 2024-10-17 16:22:11 +08:00
shanexi f6b746b213 Support number 2024-10-17 15:55:18 +08:00
wl-zhao f78f07961e update LoRA Stacker 2024-10-17 15:22:02 +08:00
wl-zhao 8311254009 update error message when missing nodes 2024-10-17 15:17:40 +08:00
wl-zhao aa185a6ff3 update 2024-10-17 14:42:36 +08:00
wl-zhao 83146c06f4 update model dependency loader 2024-10-17 14:28:12 +08:00
wl-zhao 6422a94ed4 update model dependency loader 2024-10-17 14:24:22 +08:00
wl-zhao 2a59a8a01e fix integer type 2024-10-17 11:47:44 +08:00
shanexi 370975a4ac Refactor connect 2024-10-17 11:20:20 +08:00
shanexi aa984eb533 ShellAgentPluginInputInteger add manage choices and client-side form validate 2024-10-17 10:58:21 +08:00
wl-zhao cf86159484 update requirements.txt 2024-10-17 00:31:51 +08:00
wl-zhao 1ec16a1662 update 2024-10-17 00:31:40 +08:00
shanexi 46e6b3ff91 Fix choices 2024-10-16 12:08:51 +08:00
shanexi 4e929fdc1b Remove console.log 2024-10-16 11:18:12 +08:00
shanexi cc4f0fdd79 Choices 2024-10-16 11:16:33 +08:00
shanexi d6f8d9156a Revert forceInput, comfyUI state is related somehow 2024-10-16 10:57:40 +08:00
shanexi 6905252b4d Fix api 2024-10-15 17:31:13 +08:00
shanexi da20ddb4d5 Choices UI 2024-10-15 16:55:10 +08:00
shanexi 8187880d77 Fix api 2024-10-15 16:08:22 +08:00
shanexi 1e3663f6c6 Input video 2024-10-15 15:34:26 +08:00
shanexi f52789d932 Refactor rename ext 2024-10-15 12:12:06 +08:00
shanexi d5ba41de7e Convert and connect 2024-10-15 12:06:29 +08:00
shanexi 9845c02811 Add convert to shellagent widget context menu 2024-10-15 11:41:18 +08:00
shanexi e9199161ea Force input 2024-10-15 10:18:14 +08:00
shanexi 261620ea3a Update image input 2024-10-14 17:32:47 +08:00
11 changed files with 901 additions and 101 deletions
+17 -5
View File
@@ -17,16 +17,17 @@ class ShellAgentPluginInputImage:
"required": {
"input_name": (
"STRING",
{"multiline": False, "default": "input_image"},
{"multiline": False, "default": "input_image", "forceInput": False},
),
"default_value": (
"STRING", {"image_upload": True, "default": files[0] if len(files) else ""},
# "STRING", {"image_upload": True, "default": files[0] if len(files) else ""},
sorted(files), {"image_upload": True, "forceInput": False}
),
},
"optional": {
"description": (
"STRING",
{"multiline": True, "default": ""},
{"multiline": True, "default": "", "forceInput": False},
),
}
}
@@ -48,6 +49,17 @@ class ShellAgentPluginInputImage:
"url_type": "image"
}
return schema
@classmethod
def VALIDATE_INPUTS(s, input_name, default_value, description=""):
image = default_value
if image.startswith("http"):
return True
if not folder_paths.exists_annotated_filepath(image):
return "Invalid image file: {}".format(image)
return True
def run(self, input_name, default_value=None, display_name=None, description=None):
input_dir = folder_paths.get_input_directory()
@@ -56,8 +68,8 @@ class ShellAgentPluginInputImage:
if image_path.startswith('http'):
import requests
from io import BytesIO
print("Fetching image from url: ", image)
response = requests.get(image)
print("Fetching image from url: ", image_path)
response = requests.get(image_path)
image = Image.open(BytesIO(response.content))
elif image_path.startswith('data:image/png;base64,') or image_path.startswith('data:image/jpeg;base64,') or image_path.startswith('data:image/jpg;base64,'):
import base64
+7 -3
View File
@@ -146,12 +146,16 @@ class ShellAgentPluginInputInteger:
"description": (
"STRING",
{"multiline": True, "default": ""},
)
),
"choices": (
"STRING",
{"multiline": False, "default": ""},
),
}
}
RETURN_TYPES = ("FLOAT",)
RETURN_NAMES = ("float",)
RETURN_TYPES = ("INT",)
RETURN_NAMES = ("int",)
FUNCTION = "run"
+5 -1
View File
@@ -88,8 +88,12 @@ class ShellAgentPluginInputVideo:
{"multiline": False, "default": "input_video"},
),
"default_value": (
"STRING", {"video_upload": True, "default": files[0] if len(files) else ""},
sorted(files),
{ "video_upload": True }
),
# "default_value": (
# "STRING", {"video_upload": True, "default": files[0] if len(files) else ""},
# ),
},
"optional": {
"description": (
+10 -4
View File
@@ -55,9 +55,11 @@ def schema_validator(prompt):
"outputs": {}
}
for node_id, node_info in prompt.items():
node_class_type = node_info["class_type"]
node_class_type = node_info.get("class_type")
if node_class_type is None:
raise NotImplementedError(f"Missing nodes founded, please first install the missing nodes using ComfyUI Manager")
node_cls = NODE_CLASS_MAPPINGS[node_class_type]
if hasattr(node_cls, "RELATIVE_PYTHON_MODULE") and node_cls.RELATIVE_PYTHON_MODULE == "custom_nodes.ComfyUI-ShellAgent-Plugin":
if hasattr(node_cls, "RELATIVE_PYTHON_MODULE") and node_cls.RELATIVE_PYTHON_MODULE.startswith("custom_nodes.ComfyUI-ShellAgent-Plugin"):
schema = {}
if "input_name" in node_info["inputs"]:
mode = "inputs"
@@ -122,6 +124,10 @@ async def shellagent_get_file(request):
async def shellagent_export(request):
data = await request.json()
prompt = data["prompt"]
custom_dependencies = data.get("custom_dependencies", {
"models": {},
"custom_nodes": {}
})
# extra_data = data["extra_data"]
workflow_id = str(uuid.uuid4())
@@ -137,7 +143,7 @@ async def shellagent_export(request):
try:
schemas = schema_validator(prompt)
# custom_node.json
dependency_results = resolve_dependencies(prompt)
dependency_results = resolve_dependencies(prompt, custom_dependencies)
# save_root = os.path.join(WORKFLOW_ROOT, workflow_id)
# os.makedirs(save_root, exist_ok=True)
@@ -162,6 +168,6 @@ async def shellagent_export(request):
status = 400
return_dict = {
"success": False,
"message": str(traceback.print_exc())
"message": str(traceback.format_exc()),
}
return web.json_response(return_dict, status=status)
+92 -59
View File
@@ -3,64 +3,35 @@ import subprocess
import json
import logging
from functools import partial
import re
import glob
from folder_paths import models_dir as MODELS_DIR
from folder_paths import base_path as BASE_PATH
from .utils import compute_sha256, windows_to_linux_path
from .file_upload import collect_local_file, process_local_file_path_async
ComfyUIModelLoaders = {
'VAELoader': (["vae_name"], "vae"),
'CheckpointLoader': (["ckpt_name"], "checkpoints"),
'CheckpointLoaderSimple': (["ckpt_name"], "checkpoints"),
'DiffusersLoader': (["model_path"], "diffusers"),
'unCLIPCheckpointLoader': (["ckpt_name"], "checkpoints"),
'LoraLoader': (["lora_name"], "loras"),
'LoraLoaderModelOnly': (["lora_name"], "loras"),
'ControlNetLoader': (["control_net_name"], "controlnet"),
'DiffControlNetLoader': (["control_net_name"], "controlnet"),
'UNETLoader': (["unet_name"], "unet"),
'CLIPLoader': (["clip_name"], "clip"),
'DualCLIPLoader': (["clip_name1", "clip_name2"], "clip"),
'CLIPVisionLoader': (["clip_name"], "clip_vision"),
'StyleModelLoader': (["style_model_name"], "style_models"),
'GLIGENLoader': (["gligen_name"], "gligen"),
'ImageOnlyCheckpointLoader': (["ckpt_name"], "checkpoints"),
"UpscaleModelLoader": (["model_name"], "upscale_models"),
"TripleCLIPLoader": (["clip_name1", "clip_name2", "clip_name3"], "clip"),
"HypernetworkLoader": (["hypernetwork_name"], "hypernetworks")
}
# ComfyUIFileLoaders = {
# 'VAELoader': (["vae_name"], "vae"),
# 'CheckpointLoader': (["ckpt_name"], "checkpoints"),
# 'CheckpointLoaderSimple': (["ckpt_name"], "checkpoints"),
# 'DiffusersLoader': (["model_path"], "diffusers"),
# 'unCLIPCheckpointLoader': (["ckpt_name"], "checkpoints"),
# 'LoraLoader': (["lora_name"], "loras"),
# 'LoraLoaderModelOnly': (["lora_name"], "loras"),
# 'ControlNetLoader': (["control_net_name"], "controlnet"),
# 'DiffControlNetLoader': (["control_net_name"], "controlnet"),
# 'UNETLoader': (["unet_name"], "unet"),
# 'CLIPLoader': (["clip_name"], "clip"),
# 'DualCLIPLoader': (["clip_name1", "clip_name2"], "clip"),
# 'CLIPVisionLoader': (["clip_name"], "clip_vision"),
# 'StyleModelLoader': (["style_model_name"], "style_models"),
# 'GLIGENLoader': (["gligen_name"], "gligen"),
# }
model_list_json = json.load(open(os.path.join(os.path.dirname(__file__), "model_info.json")))
model_loaders_info = json.load(open(os.path.join(os.path.dirname(__file__), "model_loader_info.json")))
node_deps_info = json.load(open(os.path.join(os.path.dirname(__file__), "node_deps_info.json")))
model_suffix = [".ckpt", ".safetensors", ".bin", ".pth", ".pt", ".onnx"]
def handle_model_info(ckpt_path):
ckpt_path = windows_to_linux_path(ckpt_path)
filename = os.path.basename(ckpt_path)
dirname = os.path.dirname(ckpt_path)
save_path = dirname.split('/', 1)[1]
save_path = os.path.dirname(os.path.relpath(ckpt_path, MODELS_DIR))
metadata_path = ckpt_path + ".json"
if os.path.isfile(metadata_path):
metadata = json.load(open(metadata_path))
model_id = metadata["id"]
else:
logging.info(f"computing sha256 of {ckpt_path}")
if not os.path.isfile(ckpt_path):
raise NotImplementedError(f"please install {ckpt_path} first!")
model_id = compute_sha256(ckpt_path)
data = {
"id": model_id,
@@ -75,7 +46,7 @@ def handle_model_info(ckpt_path):
item = {
"filename": filename,
"save_path": save_path,
"save_path": windows_to_linux_path(save_path),
"urls": urls,
}
return model_id, item
@@ -94,7 +65,7 @@ def inspect_repo_version(module_path):
['git', 'config', '--get', 'remote.origin.url'],
cwd=module_path
).strip().decode()
except subprocess.CalledProcessError:
except Exception:
return result
# Get the latest commit hash
@@ -103,7 +74,7 @@ def inspect_repo_version(module_path):
['git', 'rev-parse', 'HEAD'],
cwd=module_path
).strip().decode()
except subprocess.CalledProcessError:
except Exception:
return result
# Create and return the JSON result
@@ -114,52 +85,114 @@ def inspect_repo_version(module_path):
}
return result
def fetch_model_searcher_results(model_ids):
import requests
url = "https://shellagent.myshell.ai/models_searcher/search_urls"
headers = {
"Content-Type": "application/json"
}
data = {
"sha256": model_ids
}
def resolve_dependencies(prompt): # resolve custom nodes and models at the same time
response = requests.post(url, headers=headers, json=data)
results = [item[:10] for item in response.json()]
return results
def resolve_dependencies(prompt, custom_dependencies): # resolve custom nodes and models at the same time
from nodes import NODE_CLASS_MAPPINGS
custom_nodes = []
ckpt_paths = []
file_mapping_dict = {}
for node_id, node_info in prompt.items():
node_class_type = node_info["class_type"]
node_class_type = node_info.get("class_type")
if node_class_type is None:
raise NotImplementedError(f"Missing nodes founded, please first install the missing nodes using ComfyUI Manager")
node_cls = NODE_CLASS_MAPPINGS[node_class_type]
if hasattr(node_cls, "RELATIVE_PYTHON_MODULE"):
custom_nodes.append(node_cls.RELATIVE_PYTHON_MODULE)
if node_class_type in ComfyUIModelLoaders:
input_names, save_path = ComfyUIModelLoaders[node_class_type]
for input_name in input_names:
ckpt_path = os.path.join("models", save_path, node_info["inputs"][input_name])
ckpt_paths.append(ckpt_path)
if node_class_type in model_loaders_info:
for field_name, filename in node_info["inputs"].items():
for item in model_loaders_info[node_class_type]:
pattern = item["field_name"]
if re.match(f"^{pattern}$", field_name):
ckpt_path = os.path.join(MODELS_DIR, item["save_path"], filename)
ckpt_paths.append(ckpt_path)
else:
for field_name, filename in node_info["inputs"].items():
if type(filename) != str:
continue
is_model = False
for possible_suffix in model_suffix:
if filename.endswith(possible_suffix):
is_model = True
if is_model:
print(f"find {filename}, is_model=True")
# find possible paths
matching_files = []
# Walk through all subdirectories and files in the directory
for possible_filename in glob.glob(os.path.join(MODELS_DIR, "**", "*"), recursive=True):
if os.path.isfile(possible_filename) and possible_filename.endswith(filename):
matching_files.append(possible_filename)
print(f"matched files: {matching_files}")
if len(matching_files) == 1:
ckpt_paths.append(matching_files[0])
list(map(partial(collect_local_file, mapping_dict=file_mapping_dict), node_info["inputs"].values()))
ckpt_paths = list(set(ckpt_paths))
print("ckpt_paths:", ckpt_paths)
custom_nodes = list(set(custom_nodes))
# step 0: comfyui version
comfyui_version = inspect_repo_version("./")
comfyui_version = inspect_repo_version(BASE_PATH)
# step 1: custom nodes
custom_nodes_list = []
custom_nodes_names = []
for custom_node in custom_nodes:
try:
repo_info = inspect_repo_version(custom_node.replace(".", "/"))
repo_info = inspect_repo_version(os.path.join(BASE_PATH, custom_node.replace(".", "/")))
custom_nodes_list.append(repo_info)
if repo_info["repo"] == "":
repo_info["require_recheck"] = True
if repo_info["name"] in custom_dependencies["custom_nodes"]:
repo_info["repo"] = custom_dependencies["custom_nodes"][repo_info["name"]].get("repo", "")
repo_info["commit"] = custom_dependencies["custom_nodes"][repo_info["name"]].get("commit", "")
custom_nodes_names.append(repo_info["name"])
except:
print(f"failed to resolve repo info of {custom_node}")
for repo_name in custom_nodes_names:
if repo_name in node_deps_info:
for deps_node in node_deps_info[repo_name]:
if deps_node["name"] not in custom_nodes_names:
repo_info = inspect_repo_version(os.path.join("custom_nodes", deps_node["name"]))
deps_node["commit"] = repo_info["commit"]
custom_nodes_list.append(deps_node)
# step 2: models
models_dict = {}
missing_model_ids = []
for ckpt_path in ckpt_paths:
model_id, item = handle_model_info(ckpt_path)
models_dict[model_id] = item
if len(item["urls"]) == 0:
item["require_recheck"] = True
if model_id in custom_dependencies["models"]:
item["urls"] = custom_dependencies["models"][model_id].get("urls", [])
missing_model_ids.append(model_id)
# try to fetch from myshell model searcher
missing_model_results_myshell = fetch_model_searcher_results(missing_model_ids)
for missing_model_id, missing_model_urls in zip(missing_model_ids, missing_model_results_myshell):
if len(missing_model_urls) > 0:
models_dict[missing_model_id]["require_recheck"] = False
models_dict[missing_model_id]["urls"] = missing_model_urls
print("successfully fetch results from myshell", models_dict[missing_model_id])
# step 3: handle local files
process_local_file_path_async(file_mapping_dict, max_workers=20)
files_dict = {v[0]: {"filename": v[2], "urls": [v[1]]} for v in file_mapping_dict.values()}
dependencies = {
"models": models_dict,
"files": files_dict
}
files_dict = {v[0]: {"filename": windows_to_linux_path(os.path.relpath(v[2], BASE_PATH)), "urls": [v[1]]} for v in file_mapping_dict.values()}
results = {
"comfyui_version": comfyui_version,
+10 -6
View File
@@ -3,6 +3,7 @@ import os
import requests
import time
from concurrent.futures import ThreadPoolExecutor, as_completed
import folder_paths
from .utils import compute_sha256
@@ -45,7 +46,7 @@ def upload_file_to_myshell(local_file: str) -> str:
start_time = time.time()
ext = os.path.splitext(local_file)[1]
files = [
('file', (os.path.basename(local_file), open(local_file, 'rb'), ext_to_type[ext])),
('file', (os.path.basename(local_file), open(local_file, 'rb'), ext_to_type[ext.lower()])),
]
response = requests.request("POST", server_url, headers=headers, files=files)
if response.status_code == 200:
@@ -59,18 +60,21 @@ def upload_file_to_myshell(local_file: str) -> str:
def collect_local_file(item, mapping_dict={}):
input_dir = folder_paths.get_input_directory()
if not isinstance(item, str):
return
abspath = os.path.abspath(item)
input_abspath = os.path.join(input_dir, item)
# required file type
if os.path.isfile(item):
fpath = item
elif os.path.isfile(f"input/{item}"):
fpath = f"input/{item}"
if os.path.isfile(abspath):
fpath = abspath
elif os.path.isfile(input_abspath):
fpath = input_abspath
else:
fpath = None
if fpath is not None:
ext = os.path.splitext(fpath)[1]
if ext in ext_to_type.keys():
if ext.lower() in ext_to_type.keys():
mapping_dict[item] = fpath
return
else:
+148
View File
@@ -0,0 +1,148 @@
{
"VAELoader": [
{
"field_name": "vae_name",
"save_path": "vae"
}
],
"CheckpointLoader": [
{
"field_name": "ckpt_name",
"save_path": "checkpoints"
}
],
"CheckpointLoaderSimple": [
{
"field_name": "ckpt_name",
"save_path": "checkpoints"
}
],
"DiffusersLoader": [
{
"field_name": "model_path",
"save_path": "diffusers"
}
],
"unCLIPCheckpointLoader": [
{
"field_name": "ckpt_name",
"save_path": "checkpoints"
}
],
"LoraLoader": [
{
"field_name": "lora_name",
"save_path": "loras"
}
],
"LoraLoaderModelOnly": [
{
"field_name": "lora_name",
"save_path": "loras"
}
],
"ControlNetLoader": [
{
"field_name": "control_net_name",
"save_path": "controlnet"
}
],
"DiffControlNetLoader": [
{
"field_name": "control_net_name",
"save_path": "controlnet"
}
],
"UNETLoader": [
{
"field_name": "unet_name",
"save_path": "unet"
}
],
"CLIPLoader": [
{
"field_name": "clip_name",
"save_path": "clip"
}
],
"DualCLIPLoader": [
{
"field_name": "clip_name[1-2]",
"save_path": "clip"
}
],
"CLIPVisionLoader": [
{
"field_name": "clip_name",
"save_path": "clip_vision"
}
],
"StyleModelLoader": [
{
"field_name": "style_model_name",
"save_path": "style_models"
}
],
"GLIGENLoader": [
{
"field_name": "gligen_name",
"save_path": "gligen"
}
],
"ImageOnlyCheckpointLoader": [
{
"field_name": "ckpt_name",
"save_path": "checkpoints"
}
],
"UpscaleModelLoader": [
{
"field_name": "model_name",
"save_path": "upscale_models"
}
],
"TripleCLIPLoader": [
{
"field_name": "clip_name[1-3]",
"save_path": "clip"
}
],
"HypernetworkLoader": [
{
"field_name": "hypernetwork_name",
"save_path": "hypernetworks"
}
],
"SUPIR_model_loader_v2": [
{
"field_name": "supir_model",
"save_path": "checkpoints"
}
],
"SUPIR_model_loader_v2_clip": [
{
"field_name": "supir_model",
"save_path": "checkpoints"
}
],
"Efficient Loader": [
{
"field_name": "ckpt_name",
"save_path": "checkpoints"
},
{
"field_name": "vae_name",
"save_path": "vae"
},
{
"field_name": "lora_name",
"save_path": "loras"
}
],
"LoRA Stacker": [
{
"field_name": "lora_name_([1-9]|[1-4][0-9]|50)",
"save_path": "loras"
}
]
}
+19
View File
@@ -0,0 +1,19 @@
{
"ComfyUI-Easy-Use": [
{
"name": "ComfyUI-Inspire-Pack",
"repo": "https://github.com/ltdrdata/ComfyUI-Inspire-Pack.git",
"commit": ""
},
{
"name": "ComfyUI-Advanced-ControlNet",
"repo": "https://github.com/Kosinkadink/ComfyUI-Advanced-ControlNet.git",
"commit": ""
},
{
"name": "ComfyUI_smZNodes",
"repo": "https://github.com/shiimizu/ComfyUI_smZNodes.git",
"commit": ""
}
]
}
+6
View File
@@ -0,0 +1,6 @@
aiofiles
pydantic
opencv-python
imageio-ffmpeg
brotli
# logfire
+2 -2
View File
@@ -1,9 +1,9 @@
import hashlib
import time
from pathlib import PurePosixPath, Path
from pathlib import PurePosixPath, Path, PureWindowsPath
def windows_to_linux_path(windows_path):
return str(PurePosixPath(Path(windows_path)))
return PureWindowsPath(windows_path).as_posix()
def compute_sha256(file_path, chunk_size=1024 ** 2):
# Create a new sha256 hash object
+585 -21
View File
@@ -1,11 +1,12 @@
import { app } from "../../scripts/app.js";
import { api } from "../../scripts/api.js";
app.registerExtension({
name: "Shellagent.extension",
async setup() {
window.parent.postMessage({
type: 'loaded'
}, '*');
window.parent.postMessage({
type: 'loaded'
}, '*');
window.addEventListener('message', (event) => {
if (event.data.type === 'save') {
app.graphToPrompt().then(data => {
@@ -16,23 +17,586 @@ app.registerExtension({
}, "*");
});
}
if (event.data.type === 'load') {
app.loadGraphData(event.data.data, true, false);
}
if (event.data.type === 'load_default') {
// 使用FileReader读取JSON文件
fetch('extensions/ComfyUI-ShellAgent-Plugin/shellagent_default.json')
.then(response => response.blob())
.then(blob => {
const reader = new FileReader();
reader.onload = function(e) {
const json = JSON.parse(e.target.result);
app.loadGraphData(json, true, false);
};
reader.readAsText(blob);
})
.catch(error => console.error('加载默认JSON文件时出错:', error));
}
if (event.data.type === 'load') {
app.loadGraphData(event.data.data, true, false);
}
if (event.data.type === 'load_default') {
// 使用FileReader读取JSON文件
fetch('extensions/ComfyUI-ShellAgent-Plugin/shellagent_default.json')
.then(response => response.blob())
.then(blob => {
const reader = new FileReader();
reader.onload = function (e) {
const json = JSON.parse(e.target.result);
app.loadGraphData(json, true, false);
};
reader.readAsText(blob);
})
.catch(error => console.error('加载默认JSON文件时出错:', error));
}
});
},
});
async beforeRegisterNodeDef(nodeType, nodeData, app) {
if (["ShellAgentPluginInputText", "ShellAgentPluginInputFloat", "ShellAgentPluginInputInteger"].indexOf(nodeData.name) > -1) {
chainCallback(nodeType.prototype, "onNodeCreated", function () {
const widget = this.widgets.find(w => w.name === 'choices')
this.addWidget('button', 'manage choices', null, () => {
const container = document.createElement("div");
Object.assign(container.style, {
display: "grid",
gridTemplateColumns: "1fr 1fr",
gap: "10px",
});
const addNew = document.createElement("button");
addNew.textContent = "Add New";
addNew.classList.add("pysssss-presettext-addnew");
Object.assign(addNew.style, {
fontSize: "13px",
gridColumn: "1 / 3",
color: "dodgerblue",
width: "auto",
textAlign: "center",
});
addNew.onclick = () => {
addRow("");
};
container.append(addNew);
function addRow(p) {
const value = document.createElement("input");
if (["ShellAgentPluginInputFloat", "ShellAgentPluginInputInteger"].indexOf(nodeData.name) > -1) {
value.type = 'number';
}
const valueLbl = document.createElement("label");
value.value = p;
Object.assign(value.style, {
width: "250px",
});
valueLbl.textContent = "Value:";
valueLbl.append(value);
Object.assign(valueLbl.style, {
gridColumn: "1 / 3",
width: "auto",
});
addNew.before(valueLbl);
}
let arr = []
if (typeof widget.value === 'string') {
try {
arr = JSON.parse(widget.value)
} catch { }
} else if(Array.isArray(widget.value)) {
arr = widget.value
}
for (const a of arr) {
addRow(a);
}
const help = document.createElement("span");
help.textContent = "To remove a item set the value to blank";
help.style.gridColumn = "1 / 3";
container.append(help);
dialog.show("");
dialog.textElement.append(container);
})
const dialog = new app.ui.dialog.constructor();
dialog.element.classList.add("comfy-settings");
const closeButton = dialog.element.querySelector("button");
closeButton.textContent = "CANCEL";
const saveButton = document.createElement("button");
saveButton.textContent = "SAVE";
saveButton.onclick = function () {
const inputs = dialog.element.querySelectorAll("input");
const p = [];
for (let i = 0; i < inputs.length; i += 1) {
const v = inputs[i];
if (!v.value.trim()) {
continue;
}
p.push(v.value);
}
widget.value = p;
dialog.close();
};
closeButton.before(saveButton);
})
}
if (nodeData.name === "ShellAgentPluginInputImage") {
if (
nodeData?.input?.required?.default_value?.[1]?.image_upload === true
) {
nodeData.input.required.upload = [
"IMAGEUPLOAD",
{ widget: "default_value" },
];
}
}
if (nodeData.name === "ShellAgentPluginInputVideo") {
addUploadWidget(nodeType, nodeData, "default_value");
chainCallback(nodeType.prototype, "onNodeCreated", function () {
// const pathWidget = this.widgets.find((w) => w.name === "video");
const pathWidget = this.widgets.find((w) => w.name === "default_value");
chainCallback(pathWidget, "callback", (value) => {
if (!value) {
return;
}
let parts = ["input", value];
let extension_index = parts[1].lastIndexOf(".");
let extension = parts[1].slice(extension_index + 1);
let format = "video"
if (["gif", "webp", "avif"].includes(extension)) {
format = "image"
}
format += "/" + extension;
let params = { filename: parts[1], type: parts[0], format: format };
this.updateParameters(params, true);
});
});
addLoadVideoCommon(nodeType, nodeData);
}
if (nodeData.name.indexOf('ShellAgentPlugin') === -1) {
addMenuHandler(nodeType, function (_, options) {
if (this.widgets) {
let toInput = [];
for (const w of this.widgets) {
if (["customtext"].indexOf(w.type) > -1) {
toInput.push({
content: `${w.name} <- Input Text`,
callback: () => {
this.convertWidgetToInput(w);
const node = addNode("ShellAgentPluginInputText", this, { before: true });
const dvn = node.widgets.find(w => w.name === 'default_value')
dvn.value = w.value;
node.connect(0, this, this.inputs.length - 1);
}
})
}
if (["number"].indexOf(w.type) > -1) {
toInput.push({
content: `${w.name} <- Input Interger`,
callback: () => {
this.convertWidgetToInput(w);
const node = addNode("ShellAgentPluginInputInteger", this, { before: true });
const dvn = node.widgets.find(w => w.name === 'default_value')
dvn.value = w.value;
node.connect(0, this, this.inputs.length - 1);
}
})
toInput.push({
content: `${w.name} <- Input Float`,
callback: () => {
this.convertWidgetToInput(w);
const node = addNode("ShellAgentPluginInputFloat", this, { before: true });
const dvn = node.widgets.find(w => w.name === 'default_value')
dvn.value = w.value;
node.connect(0, this, this.inputs.length - 1);
}
})
}
}
if (toInput.length) {
options.unshift({
content: "Convert to ShellAgent",
submenu: {
options: toInput
}
})
}
}
})
}
},
});
function addMenuHandler(nodeType, cb) {
const getOpts = nodeType.prototype.getExtraMenuOptions;
nodeType.prototype.getExtraMenuOptions = function () {
const r = getOpts.apply(this, arguments);
cb.apply(this, arguments);
return r;
};
}
function fitHeight(node) {
node.setSize([node.size[0], node.computeSize([node.size[0], node.size[1]])[1]])
node?.graph?.setDirtyCanvas(true);
}
function addNode(name, nextTo, options) {
options = { select: true, shiftY: 0, before: false, ...(options || {}) };
const node = LiteGraph.createNode(name);
app.graph.add(node);
node.pos = [
options.before ? nextTo.pos[0] - node.size[0] - 30 : nextTo.pos[0] + nextTo.size[0] + 30,
nextTo.pos[1] + options.shiftY,
];
if (options.select) {
app.canvas.selectNode(node, false);
}
return node;
}
function chainCallback(object, property, callback) {
if (object == undefined) {
//This should not happen.
console.error("Tried to add callback to non-existant object")
return;
}
if (property in object && object[property]) {
const callback_orig = object[property]
object[property] = function () {
const r = callback_orig.apply(this, arguments);
callback.apply(this, arguments);
return r
};
} else {
object[property] = callback;
}
}
async function uploadFile(file) {
//TODO: Add uploaded file to cache with Cache.put()?
try {
// Wrap file in formdata so it includes filename
const body = new FormData();
const i = file.webkitRelativePath.lastIndexOf('/');
const subfolder = file.webkitRelativePath.slice(0, i + 1)
const new_file = new File([file], file.name, {
type: file.type,
lastModified: file.lastModified,
});
body.append("image", new_file);
if (i > 0) {
body.append("subfolder", subfolder);
}
const resp = await api.fetchApi("/upload/image", {
method: "POST",
body,
});
if (resp.status === 200) {
return resp
} else {
alert(resp.status + " - " + resp.statusText);
}
} catch (error) {
alert(error);
}
}
function addVideoPreview(nodeType) {
chainCallback(nodeType.prototype, "onNodeCreated", function () {
var element = document.createElement("div");
const previewNode = this;
var previewWidget = this.addDOMWidget("videopreview", "preview", element, {
serialize: false,
hideOnZoom: false,
getValue() {
return element.value;
},
setValue(v) {
element.value = v;
},
});
previewWidget.computeSize = function (width) {
if (this.aspectRatio && !this.parentEl.hidden) {
let height = (previewNode.size[0] - 20) / this.aspectRatio + 10;
if (!(height > 0)) {
height = 0;
}
this.computedHeight = height + 10;
return [width, height];
}
return [width, -4];//no loaded src, widget should not display
}
element.addEventListener('contextmenu', (e) => {
e.preventDefault()
return app.canvas._mousedown_callback(e)
}, true);
element.addEventListener('pointerdown', (e) => {
e.preventDefault()
return app.canvas._mousedown_callback(e)
}, true);
element.addEventListener('mousewheel', (e) => {
e.preventDefault()
return app.canvas._mousewheel_callback(e)
}, true);
previewWidget.value = {
hidden: false, paused: false, params: {},
muted: app.ui.settings.getSettingValue("VHS.DefaultMute", false)
}
previewWidget.parentEl = document.createElement("div");
previewWidget.parentEl.className = "vhs_preview";
previewWidget.parentEl.style['width'] = "100%"
element.appendChild(previewWidget.parentEl);
previewWidget.videoEl = document.createElement("video");
previewWidget.videoEl.controls = false;
previewWidget.videoEl.loop = true;
previewWidget.videoEl.muted = true;
previewWidget.videoEl.style['width'] = "100%"
previewWidget.videoEl.addEventListener("loadedmetadata", () => {
previewWidget.aspectRatio = previewWidget.videoEl.videoWidth / previewWidget.videoEl.videoHeight;
fitHeight(this);
});
previewWidget.videoEl.addEventListener("error", () => {
//TODO: consider a way to properly notify the user why a preview isn't shown.
previewWidget.parentEl.hidden = true;
fitHeight(this);
});
previewWidget.videoEl.onmouseenter = () => {
previewWidget.videoEl.muted = previewWidget.value.muted
};
previewWidget.videoEl.onmouseleave = () => {
previewWidget.videoEl.muted = true;
};
previewWidget.imgEl = document.createElement("img");
previewWidget.imgEl.style['width'] = "100%"
previewWidget.imgEl.hidden = true;
previewWidget.imgEl.onload = () => {
previewWidget.aspectRatio = previewWidget.imgEl.naturalWidth / previewWidget.imgEl.naturalHeight;
fitHeight(this);
};
var timeout = null;
this.updateParameters = (params, force_update) => {
if (!previewWidget.value.params) {
if (typeof (previewWidget.value != 'object')) {
previewWidget.value = { hidden: false, paused: false }
}
previewWidget.value.params = {}
}
Object.assign(previewWidget.value.params, params)
if (!force_update &&
!app.ui.settings.getSettingValue("VHS.AdvancedPreviews", false)) {
return;
}
if (timeout) {
clearTimeout(timeout);
}
if (force_update) {
previewWidget.updateSource();
} else {
timeout = setTimeout(() => previewWidget.updateSource(), 100);
}
};
previewWidget.updateSource = function () {
if (this.value.params == undefined) {
return;
}
let params = {}
Object.assign(params, this.value.params);//shallow copy
this.parentEl.hidden = this.value.hidden;
if (params.format?.split('/')[0] == 'video' ||
app.ui.settings.getSettingValue("VHS.AdvancedPreviews", false) &&
(params.format?.split('/')[1] == 'gif') || params.format == 'folder') {
this.videoEl.autoplay = !this.value.paused && !this.value.hidden;
let target_width = 256
if (element.style?.width) {
//overscale to allow scrolling. Endpoint won't return higher than native
target_width = element.style.width.slice(0, -2) * 2;
}
if (!params.force_size || params.force_size.includes("?") || params.force_size == "Disabled") {
params.force_size = target_width + "x?"
} else {
let size = params.force_size.split("x")
let ar = parseInt(size[0]) / parseInt(size[1])
params.force_size = target_width + "x" + (target_width / ar)
}
if (app.ui.settings.getSettingValue("VHS.AdvancedPreviews", false)) {
this.videoEl.src = api.apiURL('/viewvideo?' + new URLSearchParams(params));
} else {
previewWidget.videoEl.src = api.apiURL('/view?' + new URLSearchParams(params));
}
this.videoEl.hidden = false;
this.imgEl.hidden = true;
} else if (params.format?.split('/')[0] == 'image') {
//Is animated image
this.imgEl.src = api.apiURL('/view?' + new URLSearchParams(params));
this.videoEl.hidden = true;
this.imgEl.hidden = false;
}
}
previewWidget.parentEl.appendChild(previewWidget.videoEl)
previewWidget.parentEl.appendChild(previewWidget.imgEl)
});
}
function addUploadWidget(nodeType, nodeData, widgetName, type = "video") {
chainCallback(nodeType.prototype, "onNodeCreated", function () {
const pathWidget = this.widgets.find((w) => w.name === widgetName);
const fileInput = document.createElement("input");
chainCallback(this, "onRemoved", () => {
fileInput?.remove();
});
if (type == "video") {
Object.assign(fileInput, {
type: "file",
accept: "video/webm,video/mp4,video/mkv,image/gif",
style: "display: none",
onchange: async () => {
if (fileInput.files.length) {
let resp = await uploadFile(fileInput.files[0])
if (resp.status != 200) {
//upload failed and file can not be added to options
return;
}
const filename = (await resp.json()).name;
pathWidget.options.values.push(filename);
pathWidget.value = filename;
if (pathWidget.callback) {
pathWidget.callback(filename)
}
}
},
});
} else {
throw "Unknown upload type"
}
document.body.append(fileInput);
let uploadWidget = this.addWidget("button", "choose " + type + " to upload", "image", () => {
//clear the active click event
app.canvas.node_widget = null
fileInput.click();
});
uploadWidget.options.serialize = false;
});
}
function addPreviewOptions(nodeType) {
chainCallback(nodeType.prototype, "getExtraMenuOptions", function (_, options) {
// The intended way of appending options is returning a list of extra options,
// but this isn't used in widgetInputs.js and would require
// less generalization of chainCallback
let optNew = []
const previewWidget = this.widgets.find((w) => w.name === "videopreview");
let url = null
if (previewWidget.videoEl?.hidden == false && previewWidget.videoEl.src) {
//Use full quality video
url = api.apiURL('/view?' + new URLSearchParams(previewWidget.value.params));
//Workaround for 16bit png: Just do first frame
url = url.replace('%2503d', '001')
} else if (previewWidget.imgEl?.hidden == false && previewWidget.imgEl.src) {
url = previewWidget.imgEl.src;
url = new URL(url);
}
if (url) {
optNew.push(
{
content: "Open preview",
callback: () => {
window.open(url, "_blank")
},
},
{
content: "Save preview",
callback: () => {
const a = document.createElement("a");
a.href = url;
a.setAttribute("download", new URLSearchParams(previewWidget.value.params).get("filename"));
document.body.append(a);
a.click();
requestAnimationFrame(() => a.remove());
},
}
);
}
const PauseDesc = (previewWidget.value.paused ? "Resume" : "Pause") + " preview";
if (previewWidget.videoEl.hidden == false) {
optNew.push({
content: PauseDesc, callback: () => {
//animated images can't be paused and are more likely to cause performance issues.
//changing src to a single keyframe is possible,
//For now, the option is disabled if an animated image is being displayed
if (previewWidget.value.paused) {
previewWidget.videoEl?.play();
} else {
previewWidget.videoEl?.pause();
}
previewWidget.value.paused = !previewWidget.value.paused;
}
});
}
//TODO: Consider hiding elements if no video preview is available yet.
//It would reduce confusion at the cost of functionality
//(if a video preview lags the computer, the user should be able to hide in advance)
const visDesc = (previewWidget.value.hidden ? "Show" : "Hide") + " preview";
optNew.push({
content: visDesc, callback: () => {
if (!previewWidget.videoEl.hidden && !previewWidget.value.hidden) {
previewWidget.videoEl.pause();
} else if (previewWidget.value.hidden && !previewWidget.videoEl.hidden && !previewWidget.value.paused) {
previewWidget.videoEl.play();
}
previewWidget.value.hidden = !previewWidget.value.hidden;
previewWidget.parentEl.hidden = previewWidget.value.hidden;
fitHeight(this);
}
});
optNew.push({
content: "Sync preview", callback: () => {
//TODO: address case where videos have varying length
//Consider a system of sync groups which are opt-in?
for (let p of document.getElementsByClassName("vhs_preview")) {
for (let child of p.children) {
if (child.tagName == "VIDEO") {
child.currentTime = 0;
} else if (child.tagName == "IMG") {
child.src = child.src;
}
}
}
}
});
const muteDesc = (previewWidget.value.muted ? "Unmute" : "Mute") + " Preview"
optNew.push({
content: muteDesc, callback: () => {
previewWidget.value.muted = !previewWidget.value.muted
}
})
if (options.length > 0 && options[0] != null && optNew.length > 0) {
optNew.push(null);
}
options.unshift(...optNew);
});
}
function addLoadVideoCommon(nodeType, nodeData) {
addVideoPreview(nodeType);
addPreviewOptions(nodeType);
chainCallback(nodeType.prototype, "onNodeCreated", function () {
// const pathWidget = this.widgets.find((w) => w.name === "video");
const pathWidget = this.widgets.find((w) => w.name === "default_value");
//do first load
requestAnimationFrame(() => {
for (let w of [pathWidget]) {
w.callback(w.value, null, this);
}
});
});
}