diff --git a/__init__.py b/__init__.py index d0ed6d2..ea19633 100755 --- a/__init__.py +++ b/__init__.py @@ -49,6 +49,6 @@ for file in files: if hasattr(module, "NODE_DISPLAY_NAME_MAPPINGS"): NODE_DISPLAY_NAME_MAPPINGS.update(module.NODE_DISPLAY_NAME_MAPPINGS) -WEB_DIRECTORY = "web-plugin" +WEB_DIRECTORY = "web" print(NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS) __all__ = ["NODE_CLASS_MAPPINGS", "NODE_DISPLAY_NAME_MAPPINGS"] diff --git a/comfy-nodes/input_image.py b/comfy-nodes/input_image.py index a73a254..cb24158 100755 --- a/comfy-nodes/input_image.py +++ b/comfy-nodes/input_image.py @@ -37,6 +37,17 @@ class ShellAgentPluginInputImage: FUNCTION = "run" CATEGORY = "shellagent" + + @classmethod + def validate(cls, **kwargs): + schema = { + "title": kwargs["input_name"], + "type": "string", + "default": kwargs["default_value"], + "description": kwargs["description"], + "url_type": "image" + } + return schema def run(self, input_name, default_value=None, display_name=None, description=None): input_dir = folder_paths.get_input_directory() @@ -56,8 +67,9 @@ class ShellAgentPluginInputImage: decoded_image = base64.b64decode(base64_image) image = Image.open(BytesIO(decoded_image)) else: - # local path - image_path = os.path.join(input_dir, image_path) + if not os.path.isfile(image_path): # abs path + # local path + image_path = os.path.join(input_dir, image_path) image = Image.open(image_path).convert("RGB") image = ImageOps.exif_transpose(image) @@ -68,77 +80,79 @@ class ShellAgentPluginInputImage: except Exception as e: raise e -video_extensions = ["webm", "mp4", "mkv", "gif"] +# video_extensions = ["webm", "mp4", "mkv", "gif"] -class ShellAgentPluginInputVideo: - @classmethod - def INPUT_TYPES(s): - input_dir = folder_paths.get_input_directory() - files = [] - for f in os.listdir(input_dir): - if os.path.isfile(os.path.join(input_dir, f)): - file_parts = f.split(".") - if len(file_parts) > 1 and (file_parts[-1] in video_extensions): - files.append(f) +# class ShellAgentPluginInputVideo: +# @classmethod +# def INPUT_TYPES(s): +# input_dir = folder_paths.get_input_directory() +# files = [] +# for f in os.listdir(input_dir): +# if os.path.isfile(os.path.join(input_dir, f)): +# file_parts = f.split(".") +# if len(file_parts) > 1 and (file_parts[-1] in video_extensions): +# files.append(f) - return { - "required": { - "input_name": ( - "STRING", - {"multiline": False, "default": "input_video"}, - ), - "default_value": ( - "STRING", {"video_upload": True, "default": files[0] if len(files) else ""}, - ), - }, - "optional": { - "description": ( - "STRING", - {"multiline": True, "default": ""}, - ), - } - } +# return { +# "required": { +# "input_name": ( +# "STRING", +# {"multiline": False, "default": "input_video"}, +# ), +# "default_value": ( +# "STRING", {"video_upload": True, "default": files[0] if len(files) else ""}, +# ), +# }, +# "optional": { +# "description": ( +# "STRING", +# {"multiline": True, "default": ""}, +# ), +# } +# } - RETURN_TYPES = ("STRING",) - RETURN_NAMES = ("video",) +# RETURN_TYPES = ("STRING",) +# RETURN_NAMES = ("video",) - FUNCTION = "run" +# FUNCTION = "run" - CATEGORY = "shellagent" +# CATEGORY = "shellagent" - def run(self, input_name, default_value=None, description=None): - input_dir = folder_paths.get_input_directory() - if default_value.startswith("http"): - import requests +# def run(self, input_name, default_value=None, description=None): +# input_dir = folder_paths.get_input_directory() +# if default_value.startswith("http"): +# import requests - print("Fetching video from URL: ", default_value) - response = requests.get(default_value, stream=True) - file_size = int(response.headers.get("Content-Length", 0)) - file_extension = default_value.split(".")[-1].split("?")[ - 0 - ] # Extract extension and handle URLs with parameters - if file_extension not in video_extensions: - file_extension = ".mp4" +# print("Fetching video from URL: ", default_value) +# response = requests.get(default_value, stream=True) +# file_size = int(response.headers.get("Content-Length", 0)) +# file_extension = default_value.split(".")[-1].split("?")[ +# 0 +# ] # Extract extension and handle URLs with parameters +# if file_extension not in video_extensions: +# file_extension = ".mp4" - unique_filename = str(uuid.uuid4()) + "." + file_extension - video_path = os.path.join(input_dir, unique_filename) - chunk_size = 1024 # 1 Kibibyte +# unique_filename = str(uuid.uuid4()) + "." + file_extension +# video_path = os.path.join(input_dir, unique_filename) +# chunk_size = 1024 # 1 Kibibyte - num_bars = int(file_size / chunk_size) +# num_bars = int(file_size / chunk_size) - with open(video_path, "wb") as out_file: - for chunk in tqdm( - response.iter_content(chunk_size=chunk_size), - total=num_bars, - unit="KB", - desc="Downloading", - leave=True, - ): - out_file.write(chunk) - else: - video_path = os.path.abspath(os.path.join(input_dir, default_value)) +# with open(video_path, "wb") as out_file: +# for chunk in tqdm( +# response.iter_content(chunk_size=chunk_size), +# total=num_bars, +# unit="KB", +# desc="Downloading", +# leave=True, +# ): +# out_file.write(chunk) +# elif os.path.isfile(default_value): +# video_path = default_value +# else: +# video_path = os.path.abspath(os.path.join(input_dir, default_value)) - return (video_path,) +# return (video_path,) NODE_CLASS_MAPPINGS = { diff --git a/comfy-nodes/input_text.py b/comfy-nodes/input_text.py index 525a94d..fb90d7a 100755 --- a/comfy-nodes/input_text.py +++ b/comfy-nodes/input_text.py @@ -35,10 +35,159 @@ class ShellAgentPluginInputText: FUNCTION = "run" CATEGORY = "shellagent" + + @classmethod + def validate(cls, **kwargs): + schema = { + "title": kwargs["input_name"], + "type": "string", + "default": kwargs["default_value"], + "description": kwargs["description"], + } + if kwargs.get("choices", "") != "": + schema["enums"] = eval(kwargs["choices"]) + return schema def run(self, input_name, default_value=None, display_name=None, description=None, choices=None): return [default_value] + + +class ShellAgentPluginInputFloat: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "input_name": ( + "STRING", + {"multiline": False, "default": "input_float"}, + ), + }, + "optional": { + "default_value": ( + "FLOAT", + {"default": 0.}, + ), + "minimum": ( + "FLOAT", + {"default": 0.}, + ), + "maximum": ( + "FLOAT", + {"default": 0.}, + ), + "description": ( + "STRING", + {"default": ""}, + ), + "choices": ( + "STRING", + {"multiline": False, "default": ""}, + ), + } + } + + RETURN_TYPES = ("FLOAT",) + RETURN_NAMES = ("float",) + + FUNCTION = "run" + + CATEGORY = "shellagent" + + @classmethod + def validate(cls, **kwargs): + if "mininum" in kwargs and "maxinum" in kwargs and kwargs["minimum"] > kwargs["maximum"]: + raise ValueError("mininum cannot be greater than maximum") + schema = { + "title": kwargs["input_name"], + "type": "number", + "default": kwargs["default_value"], + "description": kwargs["description"], + } + if kwargs.get("choices", "") != "": + schema["enums"] = eval(kwargs["choices"]) + if "minimum" in kwargs: + schema["minimum"] = kwargs["minimum"] + if "maximum" in kwargs: + schema["maximum"] = kwargs["maximum"] + return schema -NODE_CLASS_MAPPINGS = {"ShellAgentPluginInputText": ShellAgentPluginInputText} -NODE_DISPLAY_NAME_MAPPINGS = {"ShellAgentPluginInputText": "Input Text (ShellAgent Plugin)"} \ No newline at end of file + def run(self, input_name, default_value=None, display_name=None, description=None, **kwargs): + return [default_value] + + +class ShellAgentPluginInputInteger: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "input_name": ( + "STRING", + {"multiline": False, "default": "input_integer"}, + ), + }, + "optional": { + "default_value": ( + "INT", + {"default": 0.}, + ), + "minimum": ( + "INT", + {"default": 0.}, + ), + "maximum": ( + "INT", + {"default": 0.}, + ), + "step": ( + "INT", + {"default": 1, "min": 1, "max": 10000}, + ), + "description": ( + "STRING", + {"multiline": True, "default": ""}, + ) + } + } + + RETURN_TYPES = ("FLOAT",) + RETURN_NAMES = ("float",) + + FUNCTION = "run" + + CATEGORY = "shellagent" + + @classmethod + def validate(cls, **kwargs): + if "mininum" in kwargs and "maxinum" in kwargs and kwargs["minimum"] > kwargs["maximum"]: + raise ValueError("mininum cannot be greater than maximum") + schema = { + "title": kwargs["input_name"], + "type": "integer", + "default": kwargs["default_value"], + "description": kwargs["description"], + } + if kwargs.get("choices", "") != "": + schema["enums"] = eval(kwargs["choices"]) + if "minimum" in kwargs: + schema["minimum"] = kwargs["minimum"] + if "maximum" in kwargs: + schema["maximum"] = kwargs["maximum"] + if "step" in kwargs: + schema["multiple_of"] = kwargs["step"] + return schema + + def run(self, input_name, default_value=None, display_name=None, description=None, **kwargs): + return [default_value] + + +NODE_CLASS_MAPPINGS = { + "ShellAgentPluginInputText": ShellAgentPluginInputText, + "ShellAgentPluginInputFloat": ShellAgentPluginInputFloat, + "ShellAgentPluginInputInteger": ShellAgentPluginInputInteger +} +NODE_DISPLAY_NAME_MAPPINGS = { + "ShellAgentPluginInputText": "Input Text (ShellAgent Plugin)", + "ShellAgentPluginInputFloat": "Input Float (ShellAgent Plugin)", + "ShellAgentPluginInputInteger": "Input Integer (ShellAgent Plugin)", +} \ No newline at end of file diff --git a/comfy-nodes/input_video.py b/comfy-nodes/input_video.py index a3502c0..e18d7b7 100644 --- a/comfy-nodes/input_video.py +++ b/comfy-nodes/input_video.py @@ -7,66 +7,66 @@ import uuid import tqdm -class ShellAgentPluginInputImage: - @classmethod - def INPUT_TYPES(s): - input_dir = folder_paths.get_input_directory() - files = [f for f in os.listdir(input_dir) if os.path.isfile(os.path.join(input_dir, f))] - files = sorted(files) - return { - "required": { - "input_name": ( - "STRING", - {"multiline": False, "default": "input_image"}, - ), - "default_value": ( - "STRING", {"image_upload": True, "default": files[0] if len(files) else ""}, - ), - }, - "optional": { - "description": ( - "STRING", - {"multiline": True, "default": ""}, - ), - } - } +# class ShellAgentPluginInputImage: +# @classmethod +# def INPUT_TYPES(s): +# input_dir = folder_paths.get_input_directory() +# files = [f for f in os.listdir(input_dir) if os.path.isfile(os.path.join(input_dir, f))] +# files = sorted(files) +# return { +# "required": { +# "input_name": ( +# "STRING", +# {"multiline": False, "default": "input_image"}, +# ), +# "default_value": ( +# "STRING", {"image_upload": True, "default": files[0] if len(files) else ""}, +# ), +# }, +# "optional": { +# "description": ( +# "STRING", +# {"multiline": True, "default": ""}, +# ), +# } +# } - RETURN_TYPES = ("IMAGE",) - RETURN_NAMES = ("image",) +# RETURN_TYPES = ("IMAGE",) +# RETURN_NAMES = ("image",) - FUNCTION = "run" +# FUNCTION = "run" - CATEGORY = "shellagent" +# CATEGORY = "shellagent" - def run(self, input_name, default_value=None, display_name=None, description=None): - input_dir = folder_paths.get_input_directory() - image_path = default_value - try: - if image_path.startswith('http'): - import requests - from io import BytesIO - print("Fetching image from url: ", image) - response = requests.get(image) - 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 - from io import BytesIO - print("Decoding base64 image") - base64_image = image_path[image_path.find(",")+1:] - decoded_image = base64.b64decode(base64_image) - image = Image.open(BytesIO(decoded_image)) - else: - # local path - image_path = os.path.join(input_dir, image_path) - image = Image.open(image_path).convert("RGB") +# def run(self, input_name, default_value=None, display_name=None, description=None): +# input_dir = folder_paths.get_input_directory() +# image_path = default_value +# try: +# if image_path.startswith('http'): +# import requests +# from io import BytesIO +# print("Fetching image from url: ", image) +# response = requests.get(image) +# 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 +# from io import BytesIO +# print("Decoding base64 image") +# base64_image = image_path[image_path.find(",")+1:] +# decoded_image = base64.b64decode(base64_image) +# image = Image.open(BytesIO(decoded_image)) +# else: +# # local path +# image_path = os.path.join(input_dir, image_path) +# image = Image.open(image_path).convert("RGB") - image = ImageOps.exif_transpose(image) - image = image.convert("RGB") - image = np.array(image).astype(np.float32) / 255.0 - image = torch.from_numpy(image)[None,] - return [image] - except Exception as e: - raise e +# image = ImageOps.exif_transpose(image) +# image = image.convert("RGB") +# image = np.array(image).astype(np.float32) / 255.0 +# image = torch.from_numpy(image)[None,] +# return [image] +# except Exception as e: +# raise e video_extensions = ["webm", "mp4", "mkv", "gif"] @@ -105,6 +105,17 @@ class ShellAgentPluginInputVideo: FUNCTION = "run" CATEGORY = "shellagent" + + @classmethod + def validate(cls, **kwargs): + schema = { + "title": kwargs["input_name"], + "type": "string", + "default": kwargs["default_value"], + "description": kwargs["description"], + "url_type": "video" + } + return schema def run(self, input_name, default_value=None, description=None): input_dir = folder_paths.get_input_directory() @@ -136,7 +147,10 @@ class ShellAgentPluginInputVideo: ): out_file.write(chunk) else: - video_path = os.path.abspath(os.path.join(input_dir, default_value)) + if os.path.isfile(default_value): + video_path = default_value + else: + video_path = os.path.abspath(os.path.join(input_dir, default_value)) return (video_path,) diff --git a/comfy-nodes/output_image.py b/comfy-nodes/output_image.py index 0421d2a..800a310 100644 --- a/comfy-nodes/output_image.py +++ b/comfy-nodes/output_image.py @@ -1,7 +1,8 @@ import folder_paths from nodes import SaveImage +import os -class ShellAgentSaveImage(SaveImage): +class ShellAgentSaveImages(SaveImage): @classmethod def INPUT_TYPES(s): return { @@ -15,15 +16,84 @@ class ShellAgentSaveImage(SaveImage): }, } + CATEGORY = "shellagent" + + @classmethod + def validate(cls, **kwargs): + schema = { + "title": kwargs["output_name"], + "type": "array", + "items": { + "type": "string", + "url_type": "image", + } + } + return schema + def save_images(self, images, filename_prefix="ComfyUI", prompt=None, extra_pnginfo=None, **extra_kwargs): results = super().save_images(images, filename_prefix, prompt, extra_pnginfo) results["shellagent_kwargs"] = extra_kwargs return results +class ShellAgentSaveImage(ShellAgentSaveImages): + @classmethod + def validate(cls, **kwargs): + schema = { + "title": kwargs["output_name"], + "type": "string", + "url_type": "image", + } + return schema + + +class ShellAgentSaveVideoVHS: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "filenames": ("VHS_FILENAMES", {"tooltip": "The filenames to save."}), + "output_name": ("STRING", {"multiline": False, "default": "output_video"},), + }, + } + + RETURN_TYPES = () + FUNCTION = "save_video" + + OUTPUT_NODE = True + + CATEGORY = "shellagent" + DESCRIPTION = "Saves the input images to your ComfyUI output directory." + + @classmethod + def validate(cls, **kwargs): + schema = { + "title": kwargs["output_name"], + "type": "array", + "items": { + "type": "string", + "url_type": "video", + } + } + return schema + + def save_video(self, filenames, **kwargs): + status, (preview_image, video_path) = filenames + cwd = os.getcwd() + preview_image = os.path.relpath(preview_image) + video_path = os.path.relpath(video_path) + results = {"ui": {"image": [preview_image], "video": [video_path]}} + print(results) + return results + + NODE_CLASS_MAPPINGS = { "ShellAgentPluginSaveImage": ShellAgentSaveImage, + "ShellAgentPluginSaveImages": ShellAgentSaveImages, + "ShellAgentPluginSaveVideoVHS": ShellAgentSaveVideoVHS, } NODE_DISPLAY_NAME_MAPPINGS = { - "ShellAgentPluginSaveImage": "Save Image (ShellAgent Plugin)" + "ShellAgentPluginSaveImage": "Save Image (ShellAgent Plugin)", + "ShellAgentPluginSaveImages": "Save Images (ShellAgent Plugin)", + "ShellAgentPluginSaveVideoVHS": "Save Video - VHS (ShellAgent Plugin)", } \ No newline at end of file diff --git a/custom_routes.py b/custom_routes.py index 7fd76e7..5ef78bc 100755 --- a/custom_routes.py +++ b/custom_routes.py @@ -28,31 +28,140 @@ from aiohttp import web, ClientSession, ClientError, ClientTimeout, ClientRespon import atexit from datetime import datetime import nodes +import traceback + from .dependency_checker import resolve_dependencies +WORKFLOW_ROOT = "shellagent/comfy_workflow" +CustomNodeTypeMap = { + "ShellAgentPluginInputText": "text", + "ShellAgentPluginInputInteger": "integer", + "ShellAgentPluginInputFloat": "number", + "ShellAgentPluginInputImage": "image", + "ShellAgentPluginInputVideo": "video", + "ShellAgentPluginSaveImage": "image", + "ShellAgentPluginSaveVideoVHS": "video", +} + + +def schema_validator(prompt): + from nodes import NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS + input_names = [] + output_names = [] + schemas = { + "inputs": {}, + "outputs": {} + } + for node_id, node_info in prompt.items(): + node_class_type = node_info["class_type"] + 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": + schema = {} + if "input_name" in node_info["inputs"]: + mode = "inputs" + input_name = node_info["inputs"]["input_name"] + if input_name not in input_names: + input_names.append(input_name) + else: + raise ValueError(f"Duplicated input_name found in node {NODE_DISPLAY_NAME_MAPPINGS[node_class_type]} with ID={node_id}") + # handle the schema at the same time + schema["name"] = input_name + + elif "output_name" in node_info["inputs"]: + mode = "outputs" + output_name = node_info["inputs"]["output_name"] + if output_name not in output_names: + output_names.append(output_name) + else: + raise ValueError(f"Duplicated output_name found in node {NODE_DISPLAY_NAME_MAPPINGS[node_class_type]} with ID={node_id}") + schema["name"] = output_name + else: + # neither input nor output + continue + if hasattr(node_cls, "validate"): + schema = node_cls.validate(**node_info["inputs"]) + else: + raise NotImplementedError("the validate is not implemented") + schemas[mode][node_id] = schema + return schemas + + +@server.PromptServer.instance.routes.get("/shellagent/list_workflow") # data same as queue prompt, plus workflow_name +async def shellagent_list_workflow(request): + workflow_ids = os.listdir(WORKFLOW_ROOT) + # append the metadata + data = [] + for workflow_id in workflow_ids: + + metadata_file = os.path.join(WORKFLOW_ROOT, workflow_id, "metadata.json") + metadata = json.load(open(metadata_file)) + item = { + "id": workflow_id, + "metadata": metadata + } + data.append(item) + return web.json_response(data, status=400) + +@server.PromptServer.instance.routes.post("/shellagent/get_file") # data same as queue prompt, plus workflow_name +async def shellagent_get_file(request): + data = await request.json() + assert data["filename"] in [ + "workflow_api.json", + "dependencies.json", + "metadata.json", + "extra_data.json", + "schemas.json", + ] + + data = json.load(open(os.path.join(WORKFLOW_ROOT, data["workflow_id"], data["filename"]))) + return web.json_response(data, status=400) + @server.PromptServer.instance.routes.post("/shellagent/export") # data same as queue prompt, plus workflow_name async def shellagent_export(request): data = await request.json() - client_id = data["client_id"] prompt = data["prompt"] - extra_data = data["extra_data"] - workflow_name = data["workflow_name"] + # extra_data = data["extra_data"] workflow_id = str(uuid.uuid4()) - current_time = datetime.now().strftime('%Y-%m-%d %H:%M:%S') # metadata.json - metadata = { - "name": data["workflow_name"], - "workflow_id": workflow_id, - "create_time": datetime.now().strftime('%Y-%m-%d %H:%M:%S') - } + # metadata = { + # "name": data["workflow_name"], + # "workflow_id": workflow_id, + # "create_time": datetime.now().strftime('%Y-%m-%d %H:%M:%S') + # } - # custom_node.json - resolve_dependencies(prompt) - - workflow = prompt # used during running - - # - import pdb; pdb.set_trace() \ No newline at end of file + return_dict = {} + status = 200 + try: + schemas = schema_validator(prompt) + # custom_node.json + dependency_results = resolve_dependencies(prompt) + # save_root = os.path.join(WORKFLOW_ROOT, workflow_id) + # os.makedirs(save_root, exist_ok=True) + + # fname_mapping = { + # "workflow_api.json": prompt, + # "dependencies.json": dependency_results, + # # "metadata.json": metadata, + # # "extra_data.json": extra_data, + # "schemas.json": schemas, + # } + + # for fname, dict_to_save in fname_mapping.items(): + # with open(os.path.join(save_root, fname), "w") as f: + # json.dump(dict_to_save, f, indent=2) + + return_dict = { + "success": True, + "dependencies": dependency_results, + "schemas": schemas + } + except Exception as e: + status = 400 + return_dict = { + "success": False, + "message": str(traceback.print_exc()) + } + return web.json_response(return_dict, status=status) \ No newline at end of file diff --git a/dependency_checker.py b/dependency_checker.py index ab1da37..9e05cd0 100644 --- a/dependency_checker.py +++ b/dependency_checker.py @@ -2,7 +2,10 @@ import os import subprocess import json import logging +from functools import partial + 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"), @@ -75,6 +78,12 @@ def handle_model_info(ckpt_path): def inspect_repo_version(module_path): + # Create and return the JSON result + result = { + "name": os.path.basename(module_path), + "repo": "", + "commit": "" + } # Get the remote repository URL try: remote_url = subprocess.check_output( @@ -82,7 +91,7 @@ def inspect_repo_version(module_path): cwd=module_path ).strip().decode() except subprocess.CalledProcessError: - return {"error": "Failed to get remote repository URL"} + return result # Get the latest commit hash try: @@ -91,10 +100,11 @@ def inspect_repo_version(module_path): cwd=module_path ).strip().decode() except subprocess.CalledProcessError: - return {"error": "Failed to get commit hash"} + return result # Create and return the JSON result result = { + "name": os.path.basename(module_path), "repo": remote_url, "commit": commit_hash } @@ -105,6 +115,8 @@ def resolve_dependencies(prompt): # resolve custom nodes and models at the same 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_cls = NODE_CLASS_MAPPINGS[node_class_type] @@ -115,11 +127,21 @@ def resolve_dependencies(prompt): # resolve custom nodes and models at the same for input_name in input_names: ckpt_path = os.path.join("models", save_path, node_info["inputs"][input_name]) ckpt_paths.append(ckpt_path) + list(map(partial(collect_local_file, mapping_dict=file_mapping_dict), node_info["inputs"].values())) ckpt_paths = list(set(ckpt_paths)) custom_nodes = list(set(custom_nodes)) + # step 0: comfyui version + comfyui_version = inspect_repo_version("./") + # step 1: custom nodes - custom_nodes_list = [inspect_repo_version(custom_node.replace(".", "/")) for custom_node in custom_nodes] + custom_nodes_list = [] + for custom_node in custom_nodes: + try: + repo_info = inspect_repo_version(custom_node.replace(".", "/")) + custom_nodes_list.append(repo_info) + except: + print(f"failed to resolve repo info of {custom_node}") # step 2: models models_dict = {} @@ -127,8 +149,18 @@ def resolve_dependencies(prompt): # resolve custom nodes and models at the same model_id, item = handle_model_info(ckpt_path) models_dict[model_id] = item - # step 1: handle the custom nodes version - import pdb; pdb.set_trace() - # return ckpt_pat - # # step 1: - # for class_type in \ No newline at end of file + # 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 + } + + results = { + "comfyui_version": comfyui_version, + "custom_nodes": custom_nodes_list, + "models": models_dict, + "files": files_dict, + } + return results \ No newline at end of file diff --git a/file_upload.py b/file_upload.py new file mode 100644 index 0000000..12f43f8 --- /dev/null +++ b/file_upload.py @@ -0,0 +1,97 @@ +import logging +import os +import requests +import time +from concurrent.futures import ThreadPoolExecutor, as_completed + +from .utils import compute_sha256 + +ext_to_type = { + # image + '.png': 'image/png', + '.jpg': 'image/jpeg', + '.jpeg': 'image/jpeg', + '.gif': 'image/gif', + '.bmp': 'image/bmp', + '.webp': 'image/webp', + # video + '.mp4': 'video/mp4', + '.mkv': 'video/x-matroska', + '.webm': 'video/webm', + '.avi': 'video/x-msvideo', + '.mov': 'video/quicktime', + # audio + '.mp3': 'audio/mpeg', + '.wav': 'audio/wav', + '.m4a': 'audio/mp4', +} + +def upload_file_to_myshell(local_file: str) -> str: + ''' Now we only support upload file one-by-one + ''' + MYSHELL_KEY = os.environ.get('MYSHELL_KEY', "OPENSOURCE_FIXED") + if MYSHELL_KEY is None: + raise Exception( + f"MYSHELL_KEY not found in ENV. Please set MYSHELL_KEY in settings for CDN uploading." + ) + + server_url = "https://openapi.myshell.ai/public/v1/store" + headers = { + 'x-myshell-openapi-key': MYSHELL_KEY + } + + assert os.path.isfile(local_file) + sha256sum = compute_sha256(local_file) + 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])), + ] + response = requests.request("POST", server_url, headers=headers, files=files) + if response.status_code == 200: + end_time = time.time() + logging.info(f"{local_file} uploaded, time elapsed: {end_time - start_time}") + return [sha256sum, response.json()['url'], local_file] + else: + raise Exception( + f"[HTTP ERROR] {response.status_code} - {response.text} \n" + ) + + +def collect_local_file(item, mapping_dict={}): + if not isinstance(item, str): + return + # required file type + if os.path.isfile(item): + fpath = item + elif os.path.isfile(f"input/{item}"): + fpath = f"input/{item}" + else: + fpath = None + if fpath is not None: + ext = os.path.splitext(fpath)[1] + if ext in ext_to_type.keys(): + mapping_dict[item] = fpath + return + else: + return + +def process_local_file_path_async(mapping_dict, max_workers=10): + # Using ThreadPoolExecutor for concurrent file processing + logging.info(f"upload start, {len(mapping_dict)} to upload") + start_time = time.time() + with ThreadPoolExecutor(max_workers=max_workers) as executor: + # Submit tasks to the executor + futures = {executor.submit(upload_file_to_myshell, full_path): filename for filename, full_path in mapping_dict.items()} + logging.info("submit done") + # Collect the results as they complete + for future in as_completed(futures): + filename = futures[future] + try: + result = future.result() + mapping_dict[filename] = result + except Exception as e: + print(f"Error processing {filename}: {e}") + end_time = time.time() + logging.info(f"upload end, elapsed time: {end_time - start_time}") + return \ No newline at end of file diff --git a/utils.py b/utils.py index 45f541f..3df3c03 100644 --- a/utils.py +++ b/utils.py @@ -1,5 +1,5 @@ import hashlib - +import time from pathlib import PurePosixPath, Path def windows_to_linux_path(windows_path): @@ -7,6 +7,7 @@ def windows_to_linux_path(windows_path): def compute_sha256(file_path, chunk_size=1024 ** 2): # Create a new sha256 hash object + start = time.time() sha256 = hashlib.sha256() print("start compute sha256 for", file_path) # Open the file in binary mode @@ -14,6 +15,6 @@ def compute_sha256(file_path, chunk_size=1024 ** 2): # Read the file in chunks to handle large files efficiently while chunk := file.read(chunk_size): sha256.update(chunk) - print("finish compute sha256 for", file_path) + print("finish compute sha256 for", file_path, f"time: {time.time() - start}") # Return the hexadecimal digest of the hash return sha256.hexdigest() \ No newline at end of file diff --git a/web/shellagent.js b/web/shellagent.js new file mode 100644 index 0000000..e78bf5f --- /dev/null +++ b/web/shellagent.js @@ -0,0 +1,23 @@ +import { app } from "../../scripts/app.js"; +app.registerExtension({ + name: "Shellagent.extension", + async setup() { + window.parent.postMessage({ + type: 'loaded' + }, '*'); + window.addEventListener('message', (event) => { + if (event.data.type === 'save') { + app.graphToPrompt().then(data => { + window.parent.postMessage({ + prompt: data?.output || {}, + workflow: data?.workflow || {}, + type: 'save' + }, "*"); + }); + } + if (event.data.type === 'load') { + app.loadGraphData(event.data.data, true, false); + } + }); + }, +}); \ No newline at end of file