diff --git a/nodes/attribution_by_image_node.py b/nodes/attribution_by_image_node.py index 96cc22a..29f0201 100644 --- a/nodes/attribution_by_image_node.py +++ b/nodes/attribution_by_image_node.py @@ -1,7 +1,7 @@ import requests import torch -from .common import preprocess_image, image_to_base64, poll_status_until_completed +from .common import deserialize_and_get_comfy_key, preprocess_image, image_to_base64, poll_status_until_completed class AttributionByImageNode(): @classmethod @@ -26,6 +26,7 @@ class AttributionByImageNode(): def execute(self, image, model_version, api_key): if api_key.strip() == "" or api_key.strip() == "BRIA_API_TOKEN": raise Exception("Please insert a valid API key.") + api_key = deserialize_and_get_comfy_key(api_key) # Check if image is tensor, if so, convert to NumPy array if isinstance(image, torch.Tensor): diff --git a/nodes/common.py b/nodes/common.py index c106d94..890997f 100644 --- a/nodes/common.py +++ b/nodes/common.py @@ -6,6 +6,7 @@ import base64 from torchvision.transforms import ToPILImage import requests import time +import json def postprocess_image(image): result_image = Image.open(io.BytesIO(image)) @@ -48,6 +49,7 @@ def preprocess_mask(mask): def process_request(api_url, image, mask, api_key, visual_input_content_moderation, visual_output_content_moderation): if api_key.strip() == "" or api_key.strip() == "BRIA_API_TOKEN": raise Exception("Please insert a valid API key.") + api_key = deserialize_and_get_comfy_key(api_key) # Check if image and mask are tensors, if so, convert to NumPy arrays if isinstance(image, torch.Tensor): @@ -145,3 +147,15 @@ def poll_status_until_completed(status_url, api_key, timeout=360, check_interval raise Exception(f"Error checking status: {e}") raise Exception(f"Timeout reached after {timeout} seconds") + +def deserialize_and_get_comfy_key(encoded: str): + try: + decoded = base64.b64decode(encoded).decode("utf-8") + payload = json.loads(decoded) + if (payload['type'] == "comfy"): + return payload['apiKey'] + else: + raise Exception(f"Invalid token type") + except Exception as e: + raise Exception(f"Invalid token") + diff --git a/nodes/generate_image_node_v2.py b/nodes/generate_image_node_v2.py index 9411e69..0606e05 100644 --- a/nodes/generate_image_node_v2.py +++ b/nodes/generate_image_node_v2.py @@ -2,6 +2,7 @@ import requests import torch from .common import ( + deserialize_and_get_comfy_key, postprocess_image, preprocess_image, image_to_base64, @@ -96,6 +97,7 @@ class _BaseGenerateImageNodeV2: seed, images, ) + api_token = deserialize_and_get_comfy_key(api_token) headers = {"Content-Type": "application/json", "api_token": api_token} diff --git a/nodes/generative_fill_node.py b/nodes/generative_fill_node.py index 2e6f277..80c04fe 100644 --- a/nodes/generative_fill_node.py +++ b/nodes/generative_fill_node.py @@ -4,7 +4,7 @@ from PIL import Image import io import torch -from .common import preprocess_image, preprocess_mask, image_to_base64, poll_status_until_completed +from .common import deserialize_and_get_comfy_key, preprocess_image, preprocess_mask, image_to_base64, poll_status_until_completed class GenFillNode(): @@ -39,6 +39,7 @@ class GenFillNode(): def execute(self, image, mask, prompt, api_key, seed, prompt_content_moderation, visual_input_content_moderation, visual_output_content_moderation): if api_key.strip() == "" or api_key.strip() == "BRIA_API_TOKEN": raise Exception("Please insert a valid API key.") + api_key = deserialize_and_get_comfy_key(api_key) # Check if image and mask are tensors, if so, convert to NumPy arrays if isinstance(image, torch.Tensor): diff --git a/nodes/image_expansion_node.py b/nodes/image_expansion_node.py index 481cf11..cce7866 100644 --- a/nodes/image_expansion_node.py +++ b/nodes/image_expansion_node.py @@ -4,7 +4,7 @@ from PIL import Image import io import torch -from .common import image_to_base64, preprocess_image, poll_status_until_completed +from .common import deserialize_and_get_comfy_key, image_to_base64, preprocess_image, poll_status_until_completed class ImageExpansionNode(): @@ -55,6 +55,7 @@ class ImageExpansionNode(): api_key): if api_key.strip() == "" or api_key.strip() == "BRIA_API_TOKEN": raise Exception("Please insert a valid API key.") + api_key = deserialize_and_get_comfy_key(api_key) original_image_size = [int(x.strip()) for x in original_image_size.split(",")] if original_image_size else () original_image_location = [int(x.strip()) for x in original_image_location.split(",")] if original_image_location else () canvas_size = [int(x.strip()) for x in canvas_size.split(",")] if canvas_size else () diff --git a/nodes/refine_image_node_v2.py b/nodes/refine_image_node_v2.py index a7be87f..3128c1b 100644 --- a/nodes/refine_image_node_v2.py +++ b/nodes/refine_image_node_v2.py @@ -1,5 +1,5 @@ import requests -from .common import poll_status_until_completed, postprocess_image +from .common import deserialize_and_get_comfy_key, poll_status_until_completed, postprocess_image class _BaseRefineImageNodeV2: @@ -83,7 +83,7 @@ class _BaseRefineImageNodeV2: guidance_scale, seed, ) - + api_token = deserialize_and_get_comfy_key(api_token) headers = {"Content-Type": "application/json", "api_token": api_token} try: diff --git a/nodes/reimagine_node.py b/nodes/reimagine_node.py index d398eb1..5fa6484 100644 --- a/nodes/reimagine_node.py +++ b/nodes/reimagine_node.py @@ -1,6 +1,6 @@ import requests -from .common import postprocess_image, preprocess_image, image_to_base64 +from .common import deserialize_and_get_comfy_key, postprocess_image, preprocess_image, image_to_base64 class ReimagineNode(): @@ -37,7 +37,8 @@ class ReimagineNode(): steps_num, fast, structure_ref_influence, structure_image=None, tailored_model_id=None, tailored_model_influence=None, tailored_generation_prefix=None, content_moderation=0, - ): + ): + api_key = deserialize_and_get_comfy_key(api_key) payload = { "prompt": tailored_generation_prefix + prompt, "num_results": 1, diff --git a/nodes/remove_foreground_node.py b/nodes/remove_foreground_node.py index 723a362..daa811d 100644 --- a/nodes/remove_foreground_node.py +++ b/nodes/remove_foreground_node.py @@ -4,7 +4,7 @@ from PIL import Image import io import torch -from .common import preprocess_image, image_to_base64, poll_status_until_completed +from .common import deserialize_and_get_comfy_key, preprocess_image, image_to_base64, poll_status_until_completed class RemoveForegroundNode(): @@ -34,6 +34,7 @@ class RemoveForegroundNode(): def execute(self, image, visual_input_content_moderation, visual_output_content_moderation, preserve_alpha, api_key): if api_key.strip() == "" or api_key.strip() == "BRIA_API_TOKEN": raise Exception("Please insert a valid API key.") + api_key = deserialize_and_get_comfy_key(api_key) # Check if image is tensor, if so, convert to NumPy array if isinstance(image, torch.Tensor): diff --git a/nodes/replace_bg_node.py b/nodes/replace_bg_node.py index 7b2988d..29ee118 100644 --- a/nodes/replace_bg_node.py +++ b/nodes/replace_bg_node.py @@ -4,7 +4,7 @@ from PIL import Image import io import torch -from .common import image_to_base64, preprocess_image, preprocess_mask, poll_status_until_completed +from .common import deserialize_and_get_comfy_key, image_to_base64, preprocess_image, preprocess_mask, poll_status_until_completed class ReplaceBgNode(): @@ -53,6 +53,7 @@ class ReplaceBgNode(): ref_images=None,): if api_key.strip() == "" or api_key.strip() == "BRIA_API_TOKEN": raise Exception("Please insert a valid API key.") + api_key = deserialize_and_get_comfy_key(api_key) # Check if image and mask are tensors, if so, convert to NumPy arrays if isinstance(image, torch.Tensor): diff --git a/nodes/rmbg_node.py b/nodes/rmbg_node.py index ada8207..d663b23 100644 --- a/nodes/rmbg_node.py +++ b/nodes/rmbg_node.py @@ -4,7 +4,7 @@ from PIL import Image import io import torch -from .common import preprocess_image, image_to_base64, poll_status_until_completed +from .common import deserialize_and_get_comfy_key, preprocess_image, image_to_base64, poll_status_until_completed class RmbgNode(): @classmethod @@ -34,7 +34,7 @@ class RmbgNode(): def execute(self, image, visual_input_content_moderation, visual_output_content_moderation, preserve_alpha, api_key): if api_key.strip() == "" or api_key.strip() == "BRIA_API_TOKEN": raise Exception("Please insert a valid API key.") - + api_key = deserialize_and_get_comfy_key(api_key) # Check if image is tensor, if so, convert to NumPy array if isinstance(image, torch.Tensor): image = preprocess_image(image) diff --git a/nodes/tailored_gen_node.py b/nodes/tailored_gen_node.py index 1150ad4..3bb9684 100644 --- a/nodes/tailored_gen_node.py +++ b/nodes/tailored_gen_node.py @@ -1,6 +1,6 @@ import requests -from .common import postprocess_image, preprocess_image, image_to_base64 +from .common import deserialize_and_get_comfy_key, postprocess_image, preprocess_image, image_to_base64 class TailoredGenNode(): @@ -45,6 +45,7 @@ class TailoredGenNode(): guidance_method_2=None, guidance_method_2_scale=None, guidance_method_2_image=None, content_moderation=0, ): + api_key = deserialize_and_get_comfy_key(api_key) payload = { "prompt": generation_prefix + prompt, "num_results": 1, diff --git a/nodes/tailored_model_info_node.py b/nodes/tailored_model_info_node.py index a04eaf6..29a4ad8 100644 --- a/nodes/tailored_model_info_node.py +++ b/nodes/tailored_model_info_node.py @@ -1,5 +1,5 @@ import requests - +from .common import deserialize_and_get_comfy_key class TailoredModelInfoNode(): @classmethod @@ -21,6 +21,7 @@ class TailoredModelInfoNode(): # Define the execute method as expected by ComfyUI def execute(self, model_id, api_key): + api_key = deserialize_and_get_comfy_key(api_key) response = requests.get( self.api_url + model_id, headers={"api_token": api_key} diff --git a/nodes/tailored_portrait_node.py b/nodes/tailored_portrait_node.py index 190cf4e..7a8ee75 100644 --- a/nodes/tailored_portrait_node.py +++ b/nodes/tailored_portrait_node.py @@ -4,7 +4,7 @@ from PIL import Image import io import torch -from .common import image_to_base64, preprocess_image +from .common import deserialize_and_get_comfy_key, image_to_base64, preprocess_image class TailoredPortraitNode(): @classmethod @@ -12,7 +12,7 @@ class TailoredPortraitNode(): return { "required": { "image": ("IMAGE",), # Input image from another node - "tailored_model_id": ("INT",), + "tailored_model_id": ("STRING",), # API Key input with a default value "api_key": ("STRING", {"default": "BRIA_API_TOKEN"}), # API Key input with a default value }, "optional": { @@ -34,6 +34,7 @@ class TailoredPortraitNode(): def execute(self, image, tailored_model_id, api_key, seed, tailored_model_influence, id_strength): if api_key.strip() == "" or api_key.strip() == "BRIA_API_TOKEN": raise Exception("Please insert a valid API key.") + api_key = deserialize_and_get_comfy_key(api_key) # Convert the image and mask directly to if isinstance(image, torch.Tensor): if isinstance(image, torch.Tensor): @@ -44,7 +45,7 @@ class TailoredPortraitNode(): # Prepare the API request payload payload = { "id_image_file": f"{image_base64}", - "tailored_model_id": tailored_model_id, + "tailored_model_id": int(tailored_model_id), "tailored_model_influence": tailored_model_influence, "id_strength": id_strength, "seed": seed @@ -69,7 +70,7 @@ class TailoredPortraitNode(): result_image = torch.from_numpy(result_image)[None,] return (result_image,) else: - raise Exception(f"Error: API request failed with status code {response.status_code}") + raise Exception(f"Error: API request failed with status code {response.status_code} {response.text}") except Exception as e: raise Exception(f"{e}") diff --git a/nodes/text_2_image_base_node.py b/nodes/text_2_image_base_node.py index f483f0c..130b4cf 100644 --- a/nodes/text_2_image_base_node.py +++ b/nodes/text_2_image_base_node.py @@ -1,6 +1,6 @@ import requests -from .common import postprocess_image, preprocess_image, image_to_base64 +from .common import deserialize_and_get_comfy_key, postprocess_image, preprocess_image, image_to_base64 class Text2ImageBaseNode(): @@ -48,6 +48,7 @@ class Text2ImageBaseNode(): image_prompt_mode=None, image_prompt_image=None, image_prompt_scale=None, content_moderation=0, ): + api_key = deserialize_and_get_comfy_key(api_key) payload = { "prompt": prompt, "num_results": 1, diff --git a/nodes/text_2_image_fast_node.py b/nodes/text_2_image_fast_node.py index 552720e..0a32b43 100644 --- a/nodes/text_2_image_fast_node.py +++ b/nodes/text_2_image_fast_node.py @@ -1,6 +1,6 @@ import requests -from .common import postprocess_image, preprocess_image, image_to_base64 +from .common import deserialize_and_get_comfy_key, postprocess_image, preprocess_image, image_to_base64 class Text2ImageFastNode(): @@ -45,6 +45,7 @@ class Text2ImageFastNode(): image_prompt_mode=None, image_prompt_image=None, image_prompt_scale=None, content_moderation=0, ): + api_key = deserialize_and_get_comfy_key(api_key) payload = { "prompt": prompt, "num_results": 1, diff --git a/nodes/text_2_image_hd_node.py b/nodes/text_2_image_hd_node.py index df93f6c..b813a5b 100644 --- a/nodes/text_2_image_hd_node.py +++ b/nodes/text_2_image_hd_node.py @@ -1,6 +1,6 @@ import requests -from .common import postprocess_image +from .common import deserialize_and_get_comfy_key, postprocess_image class Text2ImageHDNode(): @@ -29,12 +29,13 @@ class Text2ImageHDNode(): FUNCTION = "execute" def __init__(self): - self.api_url = "https://engine.prod.bria-api.com/v1/text-to-image/hd/2.3" #"http://0.0.0.0:5000/v1/text-to-image/hd/2.3" + self.api_url = "https://engine.prod.bria-api.com/v1/text-to-image/hd/2.2" #"http://0.0.0.0:5000/v1/text-to-image/hd/2.3" def execute( self, api_key, prompt, aspect_ratio, seed, negative_prompt, steps_num, prompt_enhancement, text_guidance_scale, medium, content_moderation=0, ): + api_key = deserialize_and_get_comfy_key(api_key) payload = { "prompt": prompt, "num_results": 1, diff --git a/nodes/utils/shot_utils.py b/nodes/utils/shot_utils.py index bcdda1a..ca0326d 100644 --- a/nodes/utils/shot_utils.py +++ b/nodes/utils/shot_utils.py @@ -1,6 +1,6 @@ import requests import torch -from ..common import postprocess_image, preprocess_image, image_to_base64 +from ..common import deserialize_and_get_comfy_key, postprocess_image, preprocess_image, image_to_base64 shot_by_text_api_url = ( "https://engine.prod.bria-api.com/v1/product/lifestyle_shot_by_text" @@ -68,6 +68,7 @@ def create_text_payload( validate_api_key(api_key) + # Process image if isinstance(image, torch.Tensor): image = preprocess_image(image) @@ -126,9 +127,11 @@ def create_image_payload(image, ref_image, api_key, placement_type, **kwargs): def make_api_request(api_url, payload, api_key, Placement_type = None): """Make API request and return processed image""" - headers = {"Content-Type": "application/json", "api_token": f"{api_key}"} + try: + api_key = deserialize_and_get_comfy_key(api_key) + headers = {"Content-Type": "application/json", "api_token": f"{api_key}"} response = requests.post(api_url, json=payload, headers=headers) if response.status_code == 200: