diff --git a/.gitignore b/.gitignore index 2ec88bb..8fa55a5 100644 --- a/.gitignore +++ b/.gitignore @@ -1,4 +1,4 @@ -__pycache__/* +__pycache__/ /examples/node_modules/* .idea diff --git a/nodes/all_nodes.py b/nodes/all_nodes.py index 660097a..2fae156 100644 --- a/nodes/all_nodes.py +++ b/nodes/all_nodes.py @@ -61,7 +61,7 @@ class ServingTextOutput: return { "required": { "serving_config": ("SERVING_CONFIG",), - "text": ("STRING", {"multiline": True, "default": ""}), + "text": ("STRING", {"multiline": True, "default": "", "forceInput": True}), }, "optional": { "chained_execution": ("SHOULD_EXECUTE",), @@ -115,6 +115,69 @@ class ServingInputText: return (default,) return (serving_config[argument],) +class ServingInputTextImage: + def __init__(self): + pass + + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "serving_config": ("SERVING_CONFIG",), + "argument": ("STRING", { + "multiline": False, + "default": "prompt" + }), + "default_prompt": ("STRING", { + "multiline": True, + "default": "" + }), + }, + "optional":{ + "default_image": ("IMAGE",) + } + } + + RETURN_TYPES = ("STRING", "IMAGE",) + FUNCTION = "out" + CATEGORY = "Serving-Toolkit" + + def convert_color(self, image): + if len(image.shape) > 2 and image.shape[2] >= 4: + return cv2.cvtColor(image, cv2.COLOR_BGRA2RGB) + return cv2.cvtColor(image, cv2.COLOR_BGR2RGB) + + def load_image(self, base64_str): + nparr = np.frombuffer(base64.b64decode(base64_str), np.uint8) + result = cv2.imdecode(nparr, cv2.IMREAD_UNCHANGED) + result = self.convert_color(result) + result = result.astype(np.float32) / 255.0 + image = torch.from_numpy(result)[None,] + return image + + def out(self, serving_config, argument, default_prompt, default_image = None): + attachment_url_key = "attachment_url_0" + if attachment_url_key not in serving_config: + if default_image is not None: + return (default_image,) + raise ValueError("No attachment found in serving_config") + + attachment_url = serving_config[attachment_url_key] + response = requests.get(attachment_url) + image = Image.open(io.BytesIO(response.content)).convert("RGB") + + # Convert PIL image to base64 string + image_file = io.BytesIO() + image.save(image_file, format='PNG') + image_file.seek(0) + base64_img = base64.b64encode(image_file.read()).decode('utf-8') + + # Use the base64 string to get the image tensor + img_out = self.load_image(base64_img) + if argument not in serving_config: + return (default_prompt, img_out) + return (serving_config[argument], img_out) + class ServingInputNumber: def __init__(self): @@ -574,6 +637,7 @@ class AlwaysExecute: NODE_CLASS_MAPPINGS = { "ServingOutput": ServingOutput, "ServingInputText": ServingInputText, + "ServingInputTextImage": ServingInputTextImage, "ServingInputNumber": ServingInputNumber, "DiscordServing": DiscordServing, "WebSocketServing": WebSocketServing, @@ -591,6 +655,7 @@ NODE_DISPLAY_NAME_MAPPINGS = { "DiscordServing": "Discord Serving", "WebSocketServing": "WebSocket Serving", "ServingInputText": "Serving Input Text", + "ServingInputTextImage": "Serving Input Text & Image", "ServingInputNumber": "Serving Input Number", "ServingInputImage": "Serving Input Image", "ServingTextOutput": "Serving Text Output", diff --git a/nodes/telegram_serving.py b/nodes/telegram_serving.py index 7636580..461cbca 100644 --- a/nodes/telegram_serving.py +++ b/nodes/telegram_serving.py @@ -19,7 +19,7 @@ class TelegramServing: self.allowed_chat_ids = None def telegram_handler(self): - @self.bot.message_handler() + @self.bot.message_handler(func=lambda message: True, content_types=['photo','text']) def handle_command(message): chat_id = str(message.chat.id) if self.allowed_chat_ids and chat_id not in self.allowed_chat_ids: @@ -27,9 +27,10 @@ class TelegramServing: f"Allowed chatids are: {self.allowed_chat_ids}, but got message from user: {message.from_user.username}, chatid: {chat_id} ! Skipping message.") return # Silently ignore messages - command_name = message.text.split()[0][1:] # Extract command name without '/' - print(f"Received command from {message.chat.id}: {message.text}") - parsed_data = parse_command_string(message.text, command_name) + text = message.caption if message.content_type == 'photo' else message.text + command_name = text.split()[0][1:] # Extract command name without '/' + print(f"Received command from {message.chat.id}: {text}") + parsed_data = parse_command_string(text, command_name) async def serve_multi_image_function(images): media_group = [] @@ -41,11 +42,11 @@ class TelegramServing: img_bytes.seek(0) media_group.append(types.InputMediaPhoto(img_bytes)) - self.bot.send_media_group(message.chat.id, media_group) + self.bot.send_media_group(message.chat.id, media_group, reply_to_message_id=message.id) def serve_image_function(image, frame_duration): image_file = tensorToImageConversion(image, frame_duration) - self.bot.send_photo(message.chat.id, image_file) + self.bot.send_photo(message.chat.id, image_file, reply_to_message_id=message.id) def is_command(command): return command == command_name @@ -55,10 +56,9 @@ class TelegramServing: parsed_data["serve_multi_image_function"] = serve_multi_image_function parsed_data["serve_text_function"] = lambda text: self.bot.reply_to(message, text) - if message.document: - file_info = self.bot.get_file(message.document.file_id) - downloaded_file = self.bot.download_file(file_info.file_path) - parsed_data["attachment_url_0"] = downloaded_file + if message.photo: + file_info = self.bot.get_file(message.photo[2].file_id) + parsed_data["attachment_url_0"] = "https://api.telegram.org/file/bot{0}/{1}".format(self.bot.token, file_info.file_path) self.data.append(parsed_data) self.data_ready.set()