diff --git a/nodes/all_nodes.py b/nodes/all_nodes.py index 660097a..966adec 100644 --- a/nodes/all_nodes.py +++ b/nodes/all_nodes.py @@ -2,7 +2,7 @@ import time import threading from .discord_client import discord_client from collections import deque -from .utils import parse_command_string, tensorToImageConversion +from .utils import parse_command_string, tensorToImageConversion, CommandRegistry import discord import asyncio import requests @@ -198,7 +198,7 @@ class DiscordServing(): self.data_ready = threading.Event() self.data = deque() self.discord_token = None - pass + self.command_registry = CommandRegistry() def discord_runner(self): discord_client.run(self.discord_token) @@ -245,6 +245,10 @@ class DiscordServing(): async def on_message(message): if message.content.startswith('!'): command_name = message.content.split()[0][1:] + + if not self.command_registry.has_command(command_name): + return + parsed_data = parse_command_string(message.content, command_name) async def serve_multi_image_function(images): @@ -307,10 +311,13 @@ class WebSocketServing(): self.ws_running = False self.websocket_url= None self.ws = None + self.command_registry = CommandRegistry() pass def on_message(self,ws,message): try: parsed = json.loads(message) + if "command_name" in parsed and not self.command_registry.has_command(parsed["command_name"]): + return self.data.append(parsed) self.data_ready.set() except Exception as e: @@ -523,12 +530,14 @@ class ServingInputImageAsLatent: return ({"samples": t},) +registered_commands = [] + class CommandPickerServing: @classmethod def INPUT_TYPES(s): return { "required": { - "serving_config": ("SERVING_CONFIG",), + "serving_config": ("SERVING_CONFIG", {"lazy": True}), "command_name": ("STRING",), "should_execute": ("SHOULD_EXECUTE", {"lazy": True}), }, @@ -539,9 +548,16 @@ class CommandPickerServing: RETURN_TYPES = () OUTPUT_NODE = True + def __init__(self): + self.commandRegistry = CommandRegistry() + def check_lazy_status(self, serving_config, command_name, should_execute=None): + if serving_config is None: + self.commandRegistry.add_command(command_name) + return ["serving_config"] + if serving_config["is_command"](command_name): - return ["should_execute"] + return ["should_execute", "serving_config"] return None def out(self, serving_config, command_name, should_execute=None): diff --git a/nodes/http_serving.py b/nodes/http_serving.py index a3525f9..3199a77 100644 --- a/nodes/http_serving.py +++ b/nodes/http_serving.py @@ -4,7 +4,7 @@ from collections import deque import json import base64 from io import BytesIO -from .utils import tensorToImageConversion +from .utils import tensorToImageConversion, CommandRegistry from PIL import Image @@ -20,6 +20,7 @@ class HTTPServing: self.output = None self.html_content = None self.path = None + self.command_registry = CommandRegistry() def http_handler(self): class RequestHandler(BaseHTTPRequestHandler): @@ -32,18 +33,25 @@ class HTTPServing: self2.end_headers() def do_POST(self2): + print(f"Received POST request: {self2.path}") content_length = int(self2.headers['Content-Length']) post_data = self2.rfile.read(content_length) data = json.loads(post_data.decode('utf-8')) self.path = self2.path + if not self.command_registry.has_command(self.path[1:]): + self2.send_response(404) + self2.send_header('Content-type', 'application/json') + self2.end_headers() + self2.wfile.write(json.dumps({"error": "Command not found"}).encode('utf-8')) + return + self.data.append(data) self.data_ready.set() self.output_ready.wait() self.output_ready.clear() response = self.output - print("Response:", response) self2.send_response(200) self2.send_header('Content-type', 'application/json') # Cors diff --git a/nodes/telegram_serving.py b/nodes/telegram_serving.py index 7636580..05dceb6 100644 --- a/nodes/telegram_serving.py +++ b/nodes/telegram_serving.py @@ -7,7 +7,7 @@ from PIL import Image import base64 import numpy as np import torch -from .utils import tensorToImageConversion, parse_command_string +from .utils import tensorToImageConversion, parse_command_string, CommandRegistry class TelegramServing: @@ -17,6 +17,7 @@ class TelegramServing: self.telegram_running = False self.bot = None self.allowed_chat_ids = None + self.command_registry = CommandRegistry() def telegram_handler(self): @self.bot.message_handler() @@ -28,6 +29,9 @@ class TelegramServing: return # Silently ignore messages command_name = message.text.split()[0][1:] # Extract command name without '/' + if not self.command_registry.has_command(command_name): + return # Silently ignore wrong commands + print(f"Received command from {message.chat.id}: {message.text}") parsed_data = parse_command_string(message.text, command_name) diff --git a/nodes/utils.py b/nodes/utils.py index eedf58c..ae65692 100644 --- a/nodes/utils.py +++ b/nodes/utils.py @@ -38,3 +38,29 @@ def tensorToImageConversion(images, duration): return img_byte_array +# This class is used to store the commands that are registered by the user. +class CommandRegistry: + _instance = None + catch_all = False + + def __new__(cls): + if cls._instance is None: + cls._instance = super().__new__(cls) + cls._instance.commands = [] + return cls._instance + + def add_command(self, command): + self.commands.append(command) + + def get_commands(self): + return self.commands + + def clear_commands(self): + self.commands.clear() + + def has_command(self, command): + return command in self.commands or self.catch_all + + def add_catch_all(self): + self.catch_all = True + diff --git a/pyproject.toml b/pyproject.toml index 6503bdb..247bbcd 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -1,9 +1,9 @@ [project] name = "comfyui-serving-toolkit" description = "This extension adds nodes that allow you to easily serve your workflow (for example using a discord bot) " -version = "1.0.0" +version = "1.0.1" license = "LICENSE" -dependencies = ["discord.py", "websocket-client", "rel", "pyTelegramBotAPI"] +dependencies = ["discord.py", "websocket-client", "rel", "pyTelegramBotAPI", "opencv-python"] [project.urls] Repository = "https://github.com/matan1905/ComfyUI-Serving-Toolkit"