This commit is contained in:
Ubuntu
2025-11-09 06:51:29 +00:00
parent c1f67dd4fb
commit 7855c4dba2
17 changed files with 54 additions and 23 deletions
+2 -1
View File
@@ -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):
+14
View File
@@ -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")
+2
View File
@@ -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}
+2 -1
View File
@@ -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):
+2 -1
View File
@@ -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 ()
+2 -2
View File
@@ -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:
+3 -2
View File
@@ -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,
+2 -1
View File
@@ -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):
+2 -1
View File
@@ -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):
+2 -2
View File
@@ -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)
+2 -1
View File
@@ -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,
+2 -1
View File
@@ -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}
+5 -4
View File
@@ -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}")
+2 -1
View File
@@ -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,
+2 -1
View File
@@ -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,
+3 -2
View File
@@ -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,
+5 -2
View File
@@ -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: