Save all commands before starting to accept them to prevent executing on wrong commands
This commit is contained in:
+20
-4
@@ -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
@@ -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
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -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
@@ -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"
|
||||
|
||||
Reference in New Issue
Block a user