From 339f26b64f852edcdca47f52491acfd8ed79b617 Mon Sep 17 00:00:00 2001 From: Ubuntu Date: Tue, 14 Apr 2026 11:58:02 +0000 Subject: [PATCH] WAI-4698:set User-Agent header --- nodes/attribution_by_image_node.py | 15 ++++---- nodes/common.py | 38 +++++-------------- nodes/fibo_edit_node.py | 10 ++--- .../fibo_edit_structured_instruction_node.py | 6 +-- nodes/generate_image_lite_node_v2.py | 10 ++--- nodes/generate_image_node_v2.py | 10 ++--- ...generate_structured_prompt_lite_node_v2.py | 8 ++-- nodes/generate_structured_prompt_node_v2.py | 8 ++-- nodes/generative_fill_node.py | 15 ++++---- nodes/image_enhance_node.py | 11 ++---- nodes/image_expansion_node.py | 6 +-- nodes/product_integrate_node.py | 10 ++--- nodes/refine_image_lite_node_v2.py | 8 ++-- nodes/refine_image_node_v2.py | 8 ++-- nodes/reimagine_node.py | 10 +++-- nodes/remove_foreground_node.py | 11 ++---- nodes/replace_bg_node.py | 6 +-- nodes/rmbg_node.py | 15 ++++---- nodes/tailored_gen_node.py | 10 +++-- nodes/tailored_model_info_node.py | 5 +-- nodes/tailored_portrait_node.py | 14 +++---- nodes/text_2_image_base_node.py | 10 +++-- nodes/text_2_image_fast_node.py | 10 +++-- nodes/text_2_image_hd_node.py | 5 +-- nodes/utils/shot_utils.py | 10 +++-- .../remove_video_background_node.py | 8 +--- .../video_nodes/video_erase_elements_node.py | 8 +--- .../video_increase_resolution_node.py | 8 +--- .../video_mask_by_key_points_node.py | 9 +---- .../video_nodes/video_mask_by_prompt_node.py | 9 +---- .../video_solid_color_background_node.py | 8 +--- nodes/video_nodes/video_utils.py | 5 ++- pyproject.toml | 2 +- 33 files changed, 139 insertions(+), 187 deletions(-) diff --git a/nodes/attribution_by_image_node.py b/nodes/attribution_by_image_node.py index bb23205..c57aadb 100644 --- a/nodes/attribution_by_image_node.py +++ b/nodes/attribution_by_image_node.py @@ -1,6 +1,12 @@ import requests -from .common import deserialize_and_get_comfy_key, image_to_base64, normalize_images_input, poll_status_until_completed, to_pil_safe +from .common import ( + bria_json_headers, + image_to_base64, + normalize_images_input, + poll_status_until_completed, + to_pil_safe, +) class AttributionByImageNode(): @classmethod @@ -24,8 +30,6 @@ class AttributionByImageNode(): def execute(self, images, 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) - images = normalize_images_input(images) batch_results = [] @@ -39,10 +43,7 @@ class AttributionByImageNode(): "model_version": model_version, } - headers = { - "Content-Type": "application/json", - "api_token": f"{api_key}" - } + headers = bria_json_headers(api_key) response = requests.post(self.api_url, json=payload, headers=headers) diff --git a/nodes/common.py b/nodes/common.py index 8fa80ae..f0a614a 100644 --- a/nodes/common.py +++ b/nodes/common.py @@ -6,16 +6,17 @@ import base64 from torchvision.transforms import ToPILImage import requests import time -import json -COMFY_KEY_ERROR = ( - "Invalid Token Type\n\n" - "The API token you’ve entered is not a ComfyUI token.\n" - "Please use the valid token from your BRIA Account API Keys page:\n" - "https://platform.bria.ai/console/account/api-keys" -) +BRIA_COMFYUI_USER_AGENT = "bria/ComfyUI" +def bria_json_headers(api_token: str) -> dict: + """Headers for JSON POST requests to Bria API.""" + return { + "Content-Type": "application/json", + "api_token": api_token, + "User-Agent": BRIA_COMFYUI_USER_AGENT, + } def postprocess_image(image): result_image = Image.open(io.BytesIO(image)) result_image = result_image.convert("RGB") @@ -86,7 +87,6 @@ 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): @@ -106,10 +106,7 @@ def process_request(api_url, image, mask, api_key, visual_input_content_moderati "visual_output_content_moderation":visual_output_content_moderation } - headers = { - "Content-Type": "application/json", - "api_token": f"{api_key}" - } + headers = bria_json_headers(api_key) try: response = requests.post(api_url, json=payload, headers=headers) @@ -161,7 +158,7 @@ def poll_status_until_completed(status_url, api_key, timeout=360, check_interval Exception: If timeout is reached or API request fails """ start_time = time.time() - headers = {"api_token": api_key} + headers = bria_json_headers(api_key) while time.time() - start_time < timeout: try: @@ -185,21 +182,6 @@ def poll_status_until_completed(status_url, api_key, timeout=360, check_interval raise Exception(f"Timeout reached after {timeout} seconds") -def deserialize_and_get_comfy_key(encoded: str) -> str: - """ - Decodes a base64-encoded JSON token and returns the ComfyUI API key. - """ - try: - decoded = base64.b64decode(encoded).decode("utf-8") - payload = json.loads(decoded) - - if payload.get("type") != "comfy": - raise Exception(COMFY_KEY_ERROR) - - return payload.get("apiKey") - - except Exception as e: - raise Exception(COMFY_KEY_ERROR) def normalize_images_input(images): """ diff --git a/nodes/fibo_edit_node.py b/nodes/fibo_edit_node.py index 107cedd..8305bf3 100644 --- a/nodes/fibo_edit_node.py +++ b/nodes/fibo_edit_node.py @@ -2,12 +2,12 @@ import requests import torch from .common import ( - deserialize_and_get_comfy_key, - postprocess_image, - preprocess_image, + bria_json_headers, image_to_base64, poll_status_until_completed, + preprocess_image, preprocess_mask, + postprocess_image, ) @@ -124,9 +124,7 @@ class FIBOEditNode: guidance_scale, seed, ) - api_token = deserialize_and_get_comfy_key(api_token) - - headers = {"Content-Type": "application/json", "api_token": api_token} + headers = bria_json_headers(api_token) try: response = requests.post(self.api_url, json=payload, headers=headers) diff --git a/nodes/fibo_edit_structured_instruction_node.py b/nodes/fibo_edit_structured_instruction_node.py index 2caf901..fa5d54c 100644 --- a/nodes/fibo_edit_structured_instruction_node.py +++ b/nodes/fibo_edit_structured_instruction_node.py @@ -1,6 +1,6 @@ import requests from .common import ( - deserialize_and_get_comfy_key, + bria_json_headers, image_to_base64, normalize_images_input, poll_status_until_completed, @@ -40,8 +40,6 @@ class FIBOEditStructuredInstructionNode: def execute(self, api_token, images, instruction): self._validate_token(api_token) - api_token = deserialize_and_get_comfy_key(api_token) - # Normalize input to list of PIL images images = normalize_images_input(images) @@ -50,7 +48,7 @@ class FIBOEditStructuredInstructionNode: for idx, pil_image in enumerate(images): try: payload = self._build_payload(pil_image, instruction) - headers = {"Content-Type": "application/json", "api_token": api_token} + headers = bria_json_headers(api_token) response = requests.post(self.api_url, json=payload, headers=headers) diff --git a/nodes/generate_image_lite_node_v2.py b/nodes/generate_image_lite_node_v2.py index 994a73b..ffe28c7 100644 --- a/nodes/generate_image_lite_node_v2.py +++ b/nodes/generate_image_lite_node_v2.py @@ -1,11 +1,11 @@ import requests import torch from .common import ( - deserialize_and_get_comfy_key, - normalize_images_input, - postprocess_image, + bria_json_headers, image_to_base64, + normalize_images_input, poll_status_until_completed, + postprocess_image, ) class GenerateImageLiteNodeV2: """Lite Image Generation Node (multi-image compatible)""" @@ -80,8 +80,6 @@ class GenerateImageLiteNodeV2: images=None, ): self._validate_token(api_token) - api_token = deserialize_and_get_comfy_key(api_token) - images_list = normalize_images_input(images) if images is not None else [None] # Structured prompts per image @@ -120,7 +118,7 @@ class GenerateImageLiteNodeV2: ref_image, ) - headers = {"Content-Type": "application/json", "api_token": api_token} + headers = bria_json_headers(api_token) response = requests.post(self.api_url, json=payload, headers=headers) if response.status_code not in (200, 202): raise Exception( diff --git a/nodes/generate_image_node_v2.py b/nodes/generate_image_node_v2.py index 481247e..53a4e91 100644 --- a/nodes/generate_image_node_v2.py +++ b/nodes/generate_image_node_v2.py @@ -2,11 +2,11 @@ import requests import torch from .common import ( - deserialize_and_get_comfy_key, - normalize_images_input, - postprocess_image, + bria_json_headers, image_to_base64, + normalize_images_input, poll_status_until_completed, + postprocess_image, ) @@ -88,8 +88,6 @@ class GenerateImageNodeV2: images=None, ): self._validate_token(api_token) - api_token = deserialize_and_get_comfy_key(api_token) - images_list = normalize_images_input(images) if images is not None else [None] # Structured prompts per image @@ -128,7 +126,7 @@ class GenerateImageNodeV2: ref_image, ) - headers = {"Content-Type": "application/json", "api_token": api_token} + headers = bria_json_headers(api_token) response = requests.post(self.api_url, json=payload, headers=headers) if response.status_code not in (200, 202): diff --git a/nodes/generate_structured_prompt_lite_node_v2.py b/nodes/generate_structured_prompt_lite_node_v2.py index c698e88..f20d53d 100644 --- a/nodes/generate_structured_prompt_lite_node_v2.py +++ b/nodes/generate_structured_prompt_lite_node_v2.py @@ -1,8 +1,8 @@ import requests from .common import ( - deserialize_and_get_comfy_key, - normalize_images_input, + bria_json_headers, image_to_base64, + normalize_images_input, poll_status_until_completed, ) @@ -45,8 +45,6 @@ class GenerateStructuredPromptLiteNodeV2: def execute(self, api_token, prompt, seed, structured_prompt, images=None): self._validate_token(api_token) - api_token = deserialize_and_get_comfy_key(api_token) - images_list = normalize_images_input(images) if images is not None else [None] # Seeds per image @@ -80,7 +78,7 @@ class GenerateStructuredPromptLiteNodeV2: image, ) - headers = {"Content-Type": "application/json", "api_token": api_token} + headers = bria_json_headers(api_token) response = requests.post(self.api_url, json=payload, headers=headers) if response.status_code not in (200, 202): raise Exception( diff --git a/nodes/generate_structured_prompt_node_v2.py b/nodes/generate_structured_prompt_node_v2.py index b39a3cb..0787d93 100644 --- a/nodes/generate_structured_prompt_node_v2.py +++ b/nodes/generate_structured_prompt_node_v2.py @@ -1,9 +1,9 @@ import requests from .common import ( - deserialize_and_get_comfy_key, - normalize_images_input, + bria_json_headers, image_to_base64, + normalize_images_input, poll_status_until_completed, ) @@ -46,8 +46,6 @@ class GenerateStructuredPromptNodeV2: def execute(self, api_token, prompt, seed, structured_prompt, images=None): self._validate_token(api_token) - api_token = deserialize_and_get_comfy_key(api_token) - images_list = normalize_images_input(images) if images is not None else [None] # Seeds per image @@ -81,7 +79,7 @@ class GenerateStructuredPromptNodeV2: image ) - headers = {"Content-Type": "application/json", "api_token": api_token} + headers = bria_json_headers(api_token) response = requests.post(self.api_url, json=payload, headers=headers) if response.status_code not in (200, 202): raise Exception( diff --git a/nodes/generative_fill_node.py b/nodes/generative_fill_node.py index 3fea6f3..8a90c85 100644 --- a/nodes/generative_fill_node.py +++ b/nodes/generative_fill_node.py @@ -4,7 +4,13 @@ from PIL import Image import io import torch -from .common import deserialize_and_get_comfy_key, preprocess_image, preprocess_mask, image_to_base64, poll_status_until_completed +from .common import ( + bria_json_headers, + image_to_base64, + poll_status_until_completed, + preprocess_image, + preprocess_mask, +) class GenFillNode(): @@ -39,8 +45,6 @@ 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): image = preprocess_image(image) @@ -64,10 +68,7 @@ class GenFillNode(): "version": 2 } - headers = { - "Content-Type": "application/json", - "api_token": f"{api_key}" - } + headers = bria_json_headers(api_key) try: # Send initial request to get status URL diff --git a/nodes/image_enhance_node.py b/nodes/image_enhance_node.py index 8821044..7c3ee92 100644 --- a/nodes/image_enhance_node.py +++ b/nodes/image_enhance_node.py @@ -5,10 +5,10 @@ from PIL import Image import torch from .common import ( - deserialize_and_get_comfy_key, + bria_json_headers, image_to_base64, normalize_images_input, - poll_status_until_completed + poll_status_until_completed, ) class ImageEnhanceNode(): @@ -51,8 +51,6 @@ class ImageEnhanceNode(): # Validate 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) - # Normalize input to list of PIL images images = normalize_images_input(images) @@ -73,10 +71,7 @@ class ImageEnhanceNode(): "preserve_alpha": preserve_alpha } - headers = { - "Content-Type": "application/json", - "api_token": f"{api_key}" - } + headers = bria_json_headers(api_key) # Send request response = requests.post(self.api_url, json=payload, headers=headers) diff --git a/nodes/image_expansion_node.py b/nodes/image_expansion_node.py index 577e847..cd3d543 100644 --- a/nodes/image_expansion_node.py +++ b/nodes/image_expansion_node.py @@ -5,7 +5,7 @@ from PIL import Image import torch from .common import ( - deserialize_and_get_comfy_key, + bria_json_headers, image_to_base64, normalize_images_input, poll_status_until_completed, @@ -59,8 +59,6 @@ class ImageExpansionNode(): ): if api_key.strip() in ("", "BRIA_API_TOKEN"): raise Exception("Please insert a valid API key.") - api_key = deserialize_and_get_comfy_key(api_key) - images = normalize_images_input(images) canvas_size = [int(x.strip()) for x in canvas_size.split(",")] if canvas_size else () original_image_size = [int(x.strip()) for x in original_image_size.split(",")] if original_image_size else () @@ -107,7 +105,7 @@ class ImageExpansionNode(): "visual_output_content_moderation": visual_output_content_moderation } - headers = {"Content-Type": "application/json", "api_token": api_key} + headers = bria_json_headers(api_key) response = requests.post(self.api_url, json=payload, headers=headers) if response.status_code not in (200, 202): raise Exception(f"API request failed with status {response.status_code}: {response.text}") diff --git a/nodes/product_integrate_node.py b/nodes/product_integrate_node.py index 938c210..0c32438 100644 --- a/nodes/product_integrate_node.py +++ b/nodes/product_integrate_node.py @@ -2,11 +2,11 @@ import requests import torch from .common import ( - deserialize_and_get_comfy_key, - postprocess_image, - preprocess_image, + bria_json_headers, image_to_base64, poll_status_until_completed, + postprocess_image, + preprocess_image, ) @@ -81,8 +81,6 @@ class ProductIntegrateNode: seed, ): self._validate_token(api_token) - api_token = deserialize_and_get_comfy_key(api_token) - # Process single scene image if isinstance(scene, torch.Tensor): processed_scene = preprocess_image(scene) @@ -99,7 +97,7 @@ class ProductIntegrateNode: seed, ) - headers = {"Content-Type": "application/json", "api_token": api_token} + headers = bria_json_headers(api_token) try: response = requests.post(self.api_url, json=payload, headers=headers) diff --git a/nodes/refine_image_lite_node_v2.py b/nodes/refine_image_lite_node_v2.py index 4d39dd6..119d7fe 100644 --- a/nodes/refine_image_lite_node_v2.py +++ b/nodes/refine_image_lite_node_v2.py @@ -1,5 +1,6 @@ import requests -from .common import deserialize_and_get_comfy_key, poll_status_until_completed, postprocess_image + +from .common import bria_json_headers, poll_status_until_completed, postprocess_image @@ -96,8 +97,7 @@ class RefineImageLiteNodeV2: guidance_scale, seed, ) - api_token = deserialize_and_get_comfy_key(api_token) - headers = {"Content-Type": "application/json", "api_token": api_token} + headers = bria_json_headers(api_token) try: response = requests.post(self.api_url, json=payload, headers=headers) @@ -130,7 +130,7 @@ class RefineImageLiteNodeV2: "seed": used_seed, } - headers = {"Content-Type": "application/json", "api_token": api_token} + headers = bria_json_headers(api_token) response = requests.post(self.generate_api_url, json=payloadForImageGenetrate, headers=headers) diff --git a/nodes/refine_image_node_v2.py b/nodes/refine_image_node_v2.py index f370706..73784dd 100644 --- a/nodes/refine_image_node_v2.py +++ b/nodes/refine_image_node_v2.py @@ -1,5 +1,6 @@ import requests -from .common import deserialize_and_get_comfy_key, poll_status_until_completed, postprocess_image + +from .common import bria_json_headers, poll_status_until_completed, postprocess_image class RefineImageNodeV2: @@ -97,8 +98,7 @@ class RefineImageNodeV2: guidance_scale, seed, ) - api_token = deserialize_and_get_comfy_key(api_token) - headers = {"Content-Type": "application/json", "api_token": api_token} + headers = bria_json_headers(api_token) try: response = requests.post(self.api_url, json=payload, headers=headers) @@ -132,7 +132,7 @@ class RefineImageNodeV2: "negative_prompt":negative_prompt } - headers = {"Content-Type": "application/json", "api_token": api_token} + headers = bria_json_headers(api_token) response = requests.post(self.generate_api_url, json=payloadForImageGenetrate, headers=headers) diff --git a/nodes/reimagine_node.py b/nodes/reimagine_node.py index 5fa6484..df3f135 100644 --- a/nodes/reimagine_node.py +++ b/nodes/reimagine_node.py @@ -1,6 +1,11 @@ import requests -from .common import deserialize_and_get_comfy_key, postprocess_image, preprocess_image, image_to_base64 +from .common import ( + bria_json_headers, + image_to_base64, + postprocess_image, + preprocess_image, +) class ReimagineNode(): @@ -38,7 +43,6 @@ class ReimagineNode(): 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, @@ -59,7 +63,7 @@ class ReimagineNode(): response = requests.post( self.api_url, json=payload, - headers={"api_token": api_key} + headers=bria_json_headers(api_key), ) if response.status_code == 200: response_dict = response.json() diff --git a/nodes/remove_foreground_node.py b/nodes/remove_foreground_node.py index a3b7a7a..acfd05c 100644 --- a/nodes/remove_foreground_node.py +++ b/nodes/remove_foreground_node.py @@ -5,10 +5,10 @@ from PIL import Image import torch from .common import ( - deserialize_and_get_comfy_key, + bria_json_headers, image_to_base64, normalize_images_input, - poll_status_until_completed + poll_status_until_completed, ) class RemoveForegroundNode(): @@ -44,8 +44,6 @@ class RemoveForegroundNode(): ): 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) - images = normalize_images_input(images) batch_results = [] @@ -60,10 +58,7 @@ class RemoveForegroundNode(): "preserve_alpha": preserve_alpha } - headers = { - "Content-Type": "application/json", - "api_token": api_key - } + headers = bria_json_headers(api_key) response = requests.post(self.api_url, json=payload, headers=headers) if response.status_code not in (200, 202): diff --git a/nodes/replace_bg_node.py b/nodes/replace_bg_node.py index 8fedbd9..44a389e 100644 --- a/nodes/replace_bg_node.py +++ b/nodes/replace_bg_node.py @@ -5,7 +5,7 @@ from PIL import Image import torch from .common import ( - deserialize_and_get_comfy_key, + bria_json_headers, image_to_base64, normalize_images_input, poll_status_until_completed, @@ -60,8 +60,6 @@ class ReplaceBgNode(): ): if api_key.strip() in ("", "BRIA_API_TOKEN"): raise Exception("Please insert a valid API key.") - api_key = deserialize_and_get_comfy_key(api_key) - images = normalize_images_input(images) # Normalize reference images @@ -96,7 +94,7 @@ class ReplaceBgNode(): "force_background_detection": force_background_detection } - headers = {"Content-Type": "application/json", "api_token": api_key} + headers = bria_json_headers(api_key) response = requests.post(self.api_url, json=payload, headers=headers) if response.status_code not in (200, 202): raise Exception(f"API request failed with status {response.status_code}: {response.text}") diff --git a/nodes/rmbg_node.py b/nodes/rmbg_node.py index d224216..c1b1ba9 100644 --- a/nodes/rmbg_node.py +++ b/nodes/rmbg_node.py @@ -4,7 +4,13 @@ import numpy as np from PIL import Image import torch -from .common import deserialize_and_get_comfy_key, image_to_base64, normalize_images_input, poll_status_until_completed, to_pil_safe +from .common import ( + bria_json_headers, + image_to_base64, + normalize_images_input, + poll_status_until_completed, + to_pil_safe, +) class RmbgNode(): @classmethod @@ -32,8 +38,6 @@ class RmbgNode(): def execute(self, images, 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) - # Normalize input to list of PIL images images = normalize_images_input(images) @@ -50,10 +54,7 @@ class RmbgNode(): "preserve_alpha": preserve_alpha } - headers = { - "Content-Type": "application/json", - "api_token": f"{api_key}" - } + headers = bria_json_headers(api_key) response = requests.post(self.api_url, json=payload, headers=headers) if response.status_code not in (200, 202): diff --git a/nodes/tailored_gen_node.py b/nodes/tailored_gen_node.py index 3bb9684..6d5a29e 100644 --- a/nodes/tailored_gen_node.py +++ b/nodes/tailored_gen_node.py @@ -1,6 +1,11 @@ import requests -from .common import deserialize_and_get_comfy_key, postprocess_image, preprocess_image, image_to_base64 +from .common import ( + bria_json_headers, + image_to_base64, + postprocess_image, + preprocess_image, +) class TailoredGenNode(): @@ -45,7 +50,6 @@ 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, @@ -74,7 +78,7 @@ class TailoredGenNode(): response = requests.post( self.api_url + model_id, json=payload, - headers={"api_token": api_key} + headers=bria_json_headers(api_key), ) if response.status_code == 200: response_dict = response.json() diff --git a/nodes/tailored_model_info_node.py b/nodes/tailored_model_info_node.py index 29a4ad8..e542903 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 +from .common import bria_json_headers class TailoredModelInfoNode(): @classmethod @@ -21,10 +21,9 @@ 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} + headers=bria_json_headers(api_key), ) if response.status_code == 200: generation_prefix = response.json()["generation_prefix"] diff --git a/nodes/tailored_portrait_node.py b/nodes/tailored_portrait_node.py index db5a9bf..4e7b74d 100644 --- a/nodes/tailored_portrait_node.py +++ b/nodes/tailored_portrait_node.py @@ -4,7 +4,12 @@ import numpy as np from PIL import Image import torch -from .common import deserialize_and_get_comfy_key, image_to_base64, normalize_images_input, to_pil_safe +from .common import ( + bria_json_headers, + image_to_base64, + normalize_images_input, + to_pil_safe, +) class TailoredPortraitNode(): @classmethod @@ -41,8 +46,6 @@ class TailoredPortraitNode(): ): 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) - # Normalize images to list of PIL images images = normalize_images_input(images) @@ -60,10 +63,7 @@ class TailoredPortraitNode(): "seed": seed } - headers = { - "Content-Type": "application/json", - "api_token": api_key - } + headers = bria_json_headers(api_key) response = requests.post(self.api_url, json=payload, headers=headers) if response.status_code != 200: diff --git a/nodes/text_2_image_base_node.py b/nodes/text_2_image_base_node.py index 130b4cf..f26d282 100644 --- a/nodes/text_2_image_base_node.py +++ b/nodes/text_2_image_base_node.py @@ -1,6 +1,11 @@ import requests -from .common import deserialize_and_get_comfy_key, postprocess_image, preprocess_image, image_to_base64 +from .common import ( + bria_json_headers, + image_to_base64, + postprocess_image, + preprocess_image, +) class Text2ImageBaseNode(): @@ -48,7 +53,6 @@ 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, @@ -84,7 +88,7 @@ class Text2ImageBaseNode(): response = requests.post( self.api_url, json=payload, - headers={"api_token": api_key} + headers=bria_json_headers(api_key), ) if response.status_code == 200: response_dict = response.json() diff --git a/nodes/text_2_image_fast_node.py b/nodes/text_2_image_fast_node.py index 0a32b43..4b80231 100644 --- a/nodes/text_2_image_fast_node.py +++ b/nodes/text_2_image_fast_node.py @@ -1,6 +1,11 @@ import requests -from .common import deserialize_and_get_comfy_key, postprocess_image, preprocess_image, image_to_base64 +from .common import ( + bria_json_headers, + image_to_base64, + postprocess_image, + preprocess_image, +) class Text2ImageFastNode(): @@ -45,7 +50,6 @@ 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, @@ -77,7 +81,7 @@ class Text2ImageFastNode(): response = requests.post( self.api_url, json=payload, - headers={"api_token": api_key} + headers=bria_json_headers(api_key), ) if response.status_code == 200: response_dict = response.json() diff --git a/nodes/text_2_image_hd_node.py b/nodes/text_2_image_hd_node.py index b813a5b..b4bc489 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 deserialize_and_get_comfy_key, postprocess_image +from .common import bria_json_headers, postprocess_image class Text2ImageHDNode(): @@ -35,7 +35,6 @@ class Text2ImageHDNode(): 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, @@ -53,7 +52,7 @@ class Text2ImageHDNode(): response = requests.post( self.api_url, json=payload, - headers={"api_token": api_key} + headers=bria_json_headers(api_key), ) if response.status_code == 200: response_dict = response.json() diff --git a/nodes/utils/shot_utils.py b/nodes/utils/shot_utils.py index ca0326d..2ff3044 100644 --- a/nodes/utils/shot_utils.py +++ b/nodes/utils/shot_utils.py @@ -1,6 +1,11 @@ import requests import torch -from ..common import deserialize_and_get_comfy_key, postprocess_image, preprocess_image, image_to_base64 +from ..common import ( + bria_json_headers, + image_to_base64, + postprocess_image, + preprocess_image, +) shot_by_text_api_url = ( "https://engine.prod.bria-api.com/v1/product/lifestyle_shot_by_text" @@ -130,8 +135,7 @@ def make_api_request(api_url, payload, api_key, Placement_type = None): try: - api_key = deserialize_and_get_comfy_key(api_key) - headers = {"Content-Type": "application/json", "api_token": f"{api_key}"} + headers = bria_json_headers(api_key) response = requests.post(api_url, json=payload, headers=headers) if response.status_code == 200: diff --git a/nodes/video_nodes/remove_video_background_node.py b/nodes/video_nodes/remove_video_background_node.py index 66e35e0..48c2dca 100644 --- a/nodes/video_nodes/remove_video_background_node.py +++ b/nodes/video_nodes/remove_video_background_node.py @@ -2,7 +2,7 @@ import os import uuid import requests import folder_paths -from ..common import deserialize_and_get_comfy_key, poll_status_until_completed +from ..common import bria_json_headers, poll_status_until_completed from .video_utils import upload_video_to_s3 class RemoveVideoBackgroundNode(): @@ -55,7 +55,6 @@ class RemoveVideoBackgroundNode(): def execute(self, api_key, video_url, preserve_audio=True, output_container_and_codec="webm_vp9",): 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) video_path = None input_video_url = "" @@ -77,10 +76,7 @@ class RemoveVideoBackgroundNode(): "output_container_and_codec": output_container_and_codec } - headers = { - "Content-Type": "application/json", - "api_token": f"{api_key}" - } + headers = bria_json_headers(api_key) response = requests.post(self.api_url, json=payload, headers=headers) diff --git a/nodes/video_nodes/video_erase_elements_node.py b/nodes/video_nodes/video_erase_elements_node.py index 2680978..f270304 100644 --- a/nodes/video_nodes/video_erase_elements_node.py +++ b/nodes/video_nodes/video_erase_elements_node.py @@ -2,7 +2,7 @@ import os import uuid import requests import folder_paths -from ..common import deserialize_and_get_comfy_key, poll_status_until_completed +from ..common import bria_json_headers, poll_status_until_completed from .video_utils import upload_video_to_s3 class VideoEraseElementsNode(): @@ -60,7 +60,6 @@ class VideoEraseElementsNode(): def execute(self, api_key, video_url, mask_url="", output_container_and_codec="mp4_h264", preserve_audio=True): 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) video_path = None if video_url and video_url.strip() != "": @@ -86,10 +85,7 @@ class VideoEraseElementsNode(): "preserve_audio": preserve_audio } - headers = { - "Content-Type": "application/json", - "api_token": f"{api_key}" - } + headers = bria_json_headers(api_key) response = requests.post(self.api_url, json=payload, headers=headers) diff --git a/nodes/video_nodes/video_increase_resolution_node.py b/nodes/video_nodes/video_increase_resolution_node.py index 46829af..6df6a0f 100644 --- a/nodes/video_nodes/video_increase_resolution_node.py +++ b/nodes/video_nodes/video_increase_resolution_node.py @@ -2,7 +2,7 @@ import os import uuid import requests import folder_paths -from ..common import deserialize_and_get_comfy_key, poll_status_until_completed +from ..common import bria_json_headers, poll_status_until_completed from .video_utils import upload_video_to_s3 class VideoIncreaseResolutionNode(): @@ -57,7 +57,6 @@ class VideoIncreaseResolutionNode(): def execute(self, api_key, video_url, desired_increase='2', output_container_and_codec="mp4_h264", preserve_audio=True): 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) video_path = None if video_url and video_url.strip() != "": @@ -83,10 +82,7 @@ class VideoIncreaseResolutionNode(): "preserve_audio": preserve_audio } - headers = { - "Content-Type": "application/json", - "api_token": f"{api_key}" - } + headers = bria_json_headers(api_key) response = requests.post(self.api_url, json=payload, headers=headers) diff --git a/nodes/video_nodes/video_mask_by_key_points_node.py b/nodes/video_nodes/video_mask_by_key_points_node.py index e83c189..1e6dcaf 100644 --- a/nodes/video_nodes/video_mask_by_key_points_node.py +++ b/nodes/video_nodes/video_mask_by_key_points_node.py @@ -2,7 +2,7 @@ import os import uuid import requests import folder_paths -from ..common import deserialize_and_get_comfy_key, poll_status_until_completed +from ..common import bria_json_headers, poll_status_until_completed from .video_utils import upload_video_to_s3 import json @@ -58,8 +58,6 @@ class VideoMaskByKeyPointsNode(): def execute(self, key_points, api_key, video_url, output_container_and_codec="mp4_h264", preserve_audio=True): 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) - try: key_points_array = json.loads(key_points) except json.JSONDecodeError as e: @@ -91,10 +89,7 @@ class VideoMaskByKeyPointsNode(): "preserve_audio": preserve_audio } - headers = { - "Content-Type": "application/json", - "api_token": f"{api_key}" - } + headers = bria_json_headers(api_key) response = requests.post(self.api_url, json=payload, headers=headers) diff --git a/nodes/video_nodes/video_mask_by_prompt_node.py b/nodes/video_nodes/video_mask_by_prompt_node.py index 00ccdb6..4c80efb 100644 --- a/nodes/video_nodes/video_mask_by_prompt_node.py +++ b/nodes/video_nodes/video_mask_by_prompt_node.py @@ -2,7 +2,7 @@ import os import uuid import requests import folder_paths -from ..common import deserialize_and_get_comfy_key, poll_status_until_completed +from ..common import bria_json_headers, poll_status_until_completed from .video_utils import upload_video_to_s3 class VideoMaskByPromptNode(): @@ -57,8 +57,6 @@ class VideoMaskByPromptNode(): def execute(self, prompt, api_key, video_url, output_container_and_codec="mp4_h264", preserve_audio=True): 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) - video_path = None if video_url and video_url.strip() != "": @@ -85,10 +83,7 @@ class VideoMaskByPromptNode(): "preserve_audio": preserve_audio } - headers = { - "Content-Type": "application/json", - "api_token": f"{api_key}" - } + headers = bria_json_headers(api_key) response = requests.post(self.api_url, json=payload, headers=headers) diff --git a/nodes/video_nodes/video_solid_color_background_node.py b/nodes/video_nodes/video_solid_color_background_node.py index 3e2e401..2742260 100644 --- a/nodes/video_nodes/video_solid_color_background_node.py +++ b/nodes/video_nodes/video_solid_color_background_node.py @@ -2,7 +2,7 @@ import os import uuid import requests import folder_paths -from ..common import deserialize_and_get_comfy_key, poll_status_until_completed +from ..common import bria_json_headers, poll_status_until_completed from .video_utils import upload_video_to_s3 class VideoSolidColorBackgroundNode(): @@ -69,7 +69,6 @@ class VideoSolidColorBackgroundNode(): def execute(self, api_key, video_url, background_color="Transparent", output_container_and_codec="webm_vp9", preserve_audio=True): 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) video_path = None if video_url and video_url.strip() != "": @@ -96,10 +95,7 @@ class VideoSolidColorBackgroundNode(): "preserve_audio": preserve_audio } - headers = { - "Content-Type": "application/json", - "api_token": f"{api_key}" - } + headers = bria_json_headers(api_key) response = requests.post(self.api_url, json=payload, headers=headers) diff --git a/nodes/video_nodes/video_utils.py b/nodes/video_nodes/video_utils.py index 1bbf755..4594102 100644 --- a/nodes/video_nodes/video_utils.py +++ b/nodes/video_nodes/video_utils.py @@ -1,11 +1,14 @@ import os import requests +from ..common import BRIA_COMFYUI_USER_AGENT + def upload_video_to_s3(video_path, filename, api_token): api_url = "https://platform.prod.bria-api.com/upload-video/anonymous/presigned-url" headers = { - "Content-Type": "application/json" + "Content-Type": "application/json", + "User-Agent": BRIA_COMFYUI_USER_AGENT, } extension = os.path.splitext(filename)[1].lower() content_type_map = { diff --git a/pyproject.toml b/pyproject.toml index fe86e95..932454c 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -1,7 +1,7 @@ [project] name = "comfyui-bria-api" description = "Custom nodes for ComfyUI using BRIA's API." -version = "2.1.15" +version = "2.1.16" license = {file = "LICENSE"} [project.urls]