Save all commands before starting to accept them to prevent executing on wrong commands

This commit is contained in:
matan
2024-09-18 11:31:11 +03:00
parent 2546e7d931
commit 7d1fe57c2a
5 changed files with 63 additions and 9 deletions
+20 -4
View File
@@ -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):
+10 -2
View File
@@ -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
+5 -1
View File
@@ -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)
+26
View File
@@ -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
+2 -2
View File
@@ -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"