diff --git a/.gitignore b/.gitignore new file mode 100644 index 0000000..bee8a64 --- /dev/null +++ b/.gitignore @@ -0,0 +1 @@ +__pycache__ diff --git a/__init__.py b/__init__.py new file mode 100755 index 0000000..6787495 --- /dev/null +++ b/__init__.py @@ -0,0 +1,54 @@ +""" +@author: MyShell +@title: comfyui-shellagent-plugin +@description: +""" +import os +import sys + +sys.path.append(os.path.join(os.path.dirname(__file__))) + +import inspect +import sys +import importlib +import subprocess +import requests +import folder_paths +from folder_paths import add_model_folder_path, get_filename_list, get_folder_paths +from tqdm import tqdm + +# from . import custom_routes +# import routes + +ag_path = os.path.join(os.path.dirname(__file__)) + +def get_python_files(path): + return [f[:-3] for f in os.listdir(path) if f.endswith(".py")] + +def append_to_sys_path(path): + if path not in sys.path: + sys.path.append(path) + +paths = ["comfy-nodes"] +files = [] + +for path in paths: + full_path = os.path.join(ag_path, path) + append_to_sys_path(full_path) + files.extend(get_python_files(full_path)) + +NODE_CLASS_MAPPINGS = {} +NODE_DISPLAY_NAME_MAPPINGS = {} + +# Import all the modules and append their mappings +for file in files: + module = importlib.import_module(file) + + if hasattr(module, "NODE_CLASS_MAPPINGS"): + NODE_CLASS_MAPPINGS.update(module.NODE_CLASS_MAPPINGS) + if hasattr(module, "NODE_DISPLAY_NAME_MAPPINGS"): + NODE_DISPLAY_NAME_MAPPINGS.update(module.NODE_DISPLAY_NAME_MAPPINGS) + +WEB_DIRECTORY = "web-plugin" +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 new file mode 100755 index 0000000..a73a254 --- /dev/null +++ b/comfy-nodes/input_image.py @@ -0,0 +1,151 @@ +import folder_paths +from PIL import Image, ImageOps +import numpy as np +import torch +import os +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": ""}, + ), + } + } + + RETURN_TYPES = ("IMAGE",) + RETURN_NAMES = ("image",) + + FUNCTION = "run" + + 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") + + 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"] + +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_TYPES = ("STRING",) + RETURN_NAMES = ("video",) + + FUNCTION = "run" + + 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 + + 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 + + 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)) + + return (video_path,) + + +NODE_CLASS_MAPPINGS = { + "ShellAgentPluginInputImage": ShellAgentPluginInputImage, + # "ShellAgentPluginInputVideo": ShellAgentPluginInputVideo, +} +NODE_DISPLAY_NAME_MAPPINGS = { + "ShellAgentPluginInputImage": "Input Image (ShellAgent Plugin)", + # "ShellAgentPluginInputVideo": "Input Video (ShellAgent Plugin)" +} \ No newline at end of file diff --git a/comfy-nodes/input_text.py b/comfy-nodes/input_text.py new file mode 100755 index 0000000..525a94d --- /dev/null +++ b/comfy-nodes/input_text.py @@ -0,0 +1,44 @@ +import folder_paths +from PIL import Image, ImageOps +import numpy as np +import torch + +class ShellAgentPluginInputText: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "input_name": ( + "STRING", + {"multiline": False, "default": "input_text"}, + ), + }, + "optional": { + "default_value": ( + "STRING", + {"multiline": True, "default": ""}, + ), + "description": ( + "STRING", + {"multiline": True, "default": ""}, + ), + "choices": ( + "STRING", + {"multiline": False, "default": ""}, + ), + } + } + + RETURN_TYPES = ("STRING",) + RETURN_NAMES = ("text",) + + FUNCTION = "run" + + CATEGORY = "shellagent" + + def run(self, input_name, default_value=None, display_name=None, description=None, choices=None): + return [default_value] + + +NODE_CLASS_MAPPINGS = {"ShellAgentPluginInputText": ShellAgentPluginInputText} +NODE_DISPLAY_NAME_MAPPINGS = {"ShellAgentPluginInputText": "Input Text (ShellAgent Plugin)"} \ No newline at end of file diff --git a/comfy-nodes/input_video.py b/comfy-nodes/input_video.py new file mode 100644 index 0000000..a3502c0 --- /dev/null +++ b/comfy-nodes/input_video.py @@ -0,0 +1,151 @@ +import folder_paths +from PIL import Image, ImageOps +import numpy as np +import torch +import os +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": ""}, + ), + } + } + + RETURN_TYPES = ("IMAGE",) + RETURN_NAMES = ("image",) + + FUNCTION = "run" + + 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") + + 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"] + +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_TYPES = ("STRING",) + RETURN_NAMES = ("video",) + + FUNCTION = "run" + + 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 + + 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 + + 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)) + + return (video_path,) + + +NODE_CLASS_MAPPINGS = { + # "ShellAgentPluginInputImage": ShellAgentPluginInputImage, + "ShellAgentPluginInputVideo": ShellAgentPluginInputVideo, +} +NODE_DISPLAY_NAME_MAPPINGS = { + # "ShellAgentPluginInputImage": "Input Image (ShellAgent Plugin)", + "ShellAgentPluginInputVideo": "Input Video (ShellAgent Plugin)" +} \ No newline at end of file diff --git a/comfy-nodes/output_image.py b/comfy-nodes/output_image.py new file mode 100644 index 0000000..0421d2a --- /dev/null +++ b/comfy-nodes/output_image.py @@ -0,0 +1,29 @@ +import folder_paths +from nodes import SaveImage + +class ShellAgentSaveImage(SaveImage): + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "images": ("IMAGE", {"tooltip": "The images to save."}), + "output_name": ("STRING", {"multiline": False, "default": "output_image"},), + "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" + }, + } + + 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 + + +NODE_CLASS_MAPPINGS = { + "ShellAgentPluginSaveImage": ShellAgentSaveImage, +} +NODE_DISPLAY_NAME_MAPPINGS = { + "ShellAgentPluginSaveImage": "Save Image (ShellAgent Plugin)" +} \ No newline at end of file diff --git a/globals.py b/globals.py new file mode 100755 index 0000000..3ee2658 --- /dev/null +++ b/globals.py @@ -0,0 +1,118 @@ +import struct +from enum import Enum +import aiohttp +from typing import List, Union, Any, Optional +from PIL import Image, ImageOps +from io import BytesIO +from pydantic import BaseModel as PydanticBaseModel + +class BaseModel(PydanticBaseModel): + class Config: + arbitrary_types_allowed = True + +class Status(Enum): + NOT_STARTED = "not-started" + RUNNING = "running" + SUCCESS = "success" + FAILED = "failed" + UPLOADING = "uploading" + +class StreamingPrompt(BaseModel): + workflow_api: Any + auth_token: str + inputs: dict[str, Union[str, bytes, Image.Image]] + running_prompt_ids: set[str] = set() + status_endpoint: Optional[str] + file_upload_endpoint: Optional[str] + +class SimplePrompt(BaseModel): + status_endpoint: Optional[str] + file_upload_endpoint: Optional[str] + + token: Optional[str] + + workflow_api: dict + status: Status = Status.NOT_STARTED + progress: set = set() + last_updated_node: Optional[str] = None, + uploading_nodes: set = set() + done: bool = False + is_realtime: bool = False, + start_time: Optional[float] = None, + +sockets = dict() +prompt_metadata: dict[str, SimplePrompt] = {} +streaming_prompt_metadata: dict[str, StreamingPrompt] = {} + +class BinaryEventTypes: + PREVIEW_IMAGE = 1 + UNENCODED_PREVIEW_IMAGE = 2 + +max_output_id_length = 24 + +async def send_image(image_data, sid=None, output_id:str = None): + max_length = max_output_id_length + output_id = output_id[:max_length] + padded_output_id = output_id.ljust(max_length, '\x00') + encoded_output_id = padded_output_id.encode('ascii', 'replace') + + image_type = image_data[0] + image = image_data[1] + max_size = image_data[2] + quality = image_data[3] + if max_size is not None: + if hasattr(Image, 'Resampling'): + resampling = Image.Resampling.BILINEAR + else: + resampling = Image.ANTIALIAS + + image = ImageOps.contain(image, (max_size, max_size), resampling) + type_num = 1 + if image_type == "JPEG": + type_num = 1 + elif image_type == "PNG": + type_num = 2 + elif image_type == "WEBP": + type_num = 3 + + bytesIO = BytesIO() + header = struct.pack(">I", type_num) + # 4 bytes for the type + bytesIO.write(header) + # 10 bytes for the output_id + position_before = bytesIO.tell() + bytesIO.write(encoded_output_id) + position_after = bytesIO.tell() + bytes_written = position_after - position_before + print(f"Bytes written: {bytes_written}") + + image.save(bytesIO, format=image_type, quality=quality, compress_level=1) + preview_bytes = bytesIO.getvalue() + await send_bytes(BinaryEventTypes.PREVIEW_IMAGE, preview_bytes, sid=sid) + +async def send_socket_catch_exception(function, message): + try: + await function(message) + except (aiohttp.ClientError, aiohttp.ClientPayloadError, ConnectionResetError) as err: + print("send error:", err) + +def encode_bytes(event, data): + if not isinstance(event, int): + raise RuntimeError(f"Binary event types must be integers, got {event}") + + packed = struct.pack(">I", event) + message = bytearray(packed) + message.extend(data) + return message + +async def send_bytes(event, data, sid=None): + message = encode_bytes(event, data) + + print("sending image to ", event, sid) + + if sid is None: + _sockets = list(sockets.values()) + for ws in _sockets: + await send_socket_catch_exception(ws.send_bytes, message) + elif sid in sockets: + await send_socket_catch_exception(sockets[sid].send_bytes, message) \ No newline at end of file