Fix Telegram photo message and add ServingInputTextImage

This commit is contained in:
H.D.Tài
2024-09-16 20:46:08 +07:00
parent 2546e7d931
commit 7096da8540
3 changed files with 77 additions and 12 deletions
+1 -1
View File
@@ -1,4 +1,4 @@
__pycache__/*
__pycache__/
/examples/node_modules/*
.idea
+66 -1
View File
@@ -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",
+10 -10
View File
@@ -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()